TinyQuery-140M / tinyquery /test_model.py
karmx's picture
Release TinyQuery 139.7M from scratch with frozen weights, reproducible Mac evaluations and runtime source
b296ad4 verified
Raw
History Blame Contribute Delete
1.99 kB
"""Regression checks for causal decoding and the learned source-copy distribution."""
import unittest
import torch
from tinyquery.model import Config,TinyQuery
class ModelChecks(unittest.TestCase):
def setUp(self):
torch.set_num_threads(2); torch.manual_seed(51)
self.model=TinyQuery(Config(vocab_size=64,width=64,layers=2,heads=4,kv_heads=2,hidden=128,context=64,copy_dim=32)).eval()
def test_causality_and_cache(self):
x=torch.randint(5,64,(2,15));x[:,7]=3
original=self.model(x)[0]; changed=x.clone();changed[:,9:]=torch.randint(5,64,(2,6))
self.assertTrue(torch.allclose(original[:,:9],self.model(changed)[0][:,:9],atol=1e-6))
past=None; parts=[]
for i in range(x.shape[1]):
logits,past,_=self.model(x[:,i:i+1],past=past,use_cache=True,last_only=True);parts.append(logits)
self.assertLess(float((original-torch.cat(parts,dim=1)).abs().max().detach()),2e-5)
self.assertTrue(torch.allclose(original.exp().sum(-1),torch.ones_like(original[:,:,0]),atol=1e-6))
def test_copy_cannot_recycle_its_response(self):
with torch.no_grad():
self.model.copy_gate.weight.zero_();self.model.copy_gate.bias.fill_(-30)
self.model.copy_query.weight.zero_();self.model.copy_key.weight.zero_()
probabilities=self.model(torch.tensor([[1,5,8,3,20,20]]),last_only=True)[0].exp()[0,0]
self.assertLess(float(probabilities[20]),1e-8)
for token in [1,5,8]:self.assertAlmostEqual(float(probabilities[token]),1/3,places=6)
def test_copy_gradient(self):
x=torch.randint(5,64,(2,15));x[:,7]=3
y=torch.roll(x,shifts=-1,dims=1);mask=torch.arange(15)[None,:].expand(2,-1)>=7
loss,_=self.model(x,y,mask,torch.tensor([7,7]),torch.tensor([0,0]))
loss.backward()
self.assertTrue(torch.isfinite(loss))
self.assertGreater(float(self.model.copy_query.weight.grad.norm()),0)
if __name__=='__main__':unittest.main()