| """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() |
|
|