"""Focused invariant tests; requires torch and transformers but no downloads.""" import types import unittest import torch from torch import nn from v4_model import JudgmentV4, normalize_row class Tokenizer: pad_token_id=0 def encode(self,text,add_special_tokens=False):return [int(text)+1] if int(text)<10 else [1,2] def apply_chat_template(self,messages,**kwargs): assert kwargs.get('return_dict') is False assert kwargs.get('enable_thinking') is False # Make a predictable sequence without fetching a tokenizer. return [1,2,3] if 'short' in messages[1]['content'] else [1,2,3,4,5] class Backbone(nn.Module): def __init__(self): super().__init__();self.emb=nn.Embedding(300,8) def forward(self,input_ids,**kwargs):return types.SimpleNamespace(last_hidden_state=self.emb(input_ids)) class LM(nn.Module): def __init__(self): super().__init__();self.model=Backbone();self.head=nn.Linear(8,300,bias=False) self.config=types.SimpleNamespace(max_position_embeddings=65536,use_cache=True) def get_output_embeddings(self):return self.head def forward(self,*a,**k):raise AssertionError('Full sequence vocabulary forward must never run') class Tests(unittest.TestCase): def setUp(self):self.model=JudgmentV4(device='cpu',tokenizer=Tokenizer(),lm=LM()) def row(self,state='short',n=2):return dict(state=state,kind='choice',instructions='choose',candidates=[str(i) for i in range(n)],keys=[str(i) for i in range(n)],target=[1.]+[0.]*(n-1)) def test_noul(self): r=self.row();r.update(kind='noul',target=[.8],candidates=['']) out=normalize_row(r);self.assertEqual(out['keys'],['false','true']);self.assertAlmostEqual(out['target'][0],.2) def test_order_preserved(self): r=self.row(n=3);r['kind']='score';self.assertEqual(normalize_row(r)['candidates'],['0','1','2']) def test_no_silent_truncation(self): with self.assertRaises(ValueError):self.model.encode_rows([self.row('long')],max_length=4) self.assertEqual(self.model.encode_rows([self.row('long')],max_length=4,strict=False),[]) def test_forward_masks_gradients_padding(self): e=self.model.encode_rows([self.row(),self.row('long',3)]) b=self.model.collate(e) self.assertEqual(b['attention_mask'][0].tolist(),[0,0,1,1,1]) out=self.model(**b);self.assertTrue(torch.isfinite(out['loss'])) self.assertEqual(out['logits'].softmax(-1)[0,2].item(),0) single=self.model(**self.model.collate(e[:1]))['logits'][0] torch.testing.assert_close(single,out['logits'][0,:2]) out['loss'].backward();self.assertIsNotNone(self.model.lm.model.emb.weight.grad) def test_temperature_inference_only(self): row=self.row() self.model.eval() encoded=self.model.encode_rows([row]);batch=self.model.collate(encoded) raw=self.model(**batch)['logits'].detach().clone() self.model.calibration={'choice':{'temperature':2.}} actual=self.model.predict_records([row])[0] expected=(raw/2).softmax(-1)[0].tolist() self.assertEqual(actual,expected) torch.testing.assert_close(raw,self.model(**batch)['logits']) def test_multiple_rows_preserve_order_and_calibration(self): rows=[self.row('long',3),self.row('short',2),self.row('long',4)] rows[1]['kind']='score' self.model.calibration={'choice':{'temperature':2.},'score':{'temperature':.5}} single=[self.model.predict_records([r])[0] for r in rows] batch=self.model.predict_records(rows) self.assertEqual([len(x) for x in batch],[3,2,4]) for a,b in zip(single,batch):torch.testing.assert_close(torch.tensor(a),torch.tensor(b)) self.assertEqual(self.model.last_usage['encoded_questions'],3) def test_multitoken_indices_rejected(self): self.assertEqual(self.model.max_candidates,10) self.model.encode_rows([self.row(n=10)]) with self.assertRaisesRegex(ValueError,'at most 10'): self.model.encode_rows([self.row(n=11)]) def test_native_head_final_position_only(self): calls=[] hook=self.model.lm.head.register_forward_hook(lambda module,args,output:calls.append(args[0].shape)) encoded=self.model.encode_rows([self.row(),self.row('long',3)]) out=self.model(**self.model.collate(encoded)) self.assertEqual(calls,[torch.Size([2,8])]) hook.remove() self.assertEqual(out['logits'].shape,(2,3)) def test_multimodal_decoder_bypassed(self): decoder=self.model.lm.model class Multimodal(nn.Module): def __init__(self):super().__init__();self.language_model=decoder def forward(self,*args,**kwargs):raise AssertionError('Must use text decoder') self.model.lm.model=Multimodal() self.model(**self.model.collate(self.model.encode_rows([self.row()]))) def test_lora_includes_hybrid_excludes_vision(self): from v4_model import lora_targets lm=nn.Module();lm.model=nn.Module() lm.model.language_model=nn.Module();lm.model.visual=nn.Module() for name in ('in_proj_qkv','in_proj_z','in_proj_a','in_proj_b','out_proj','q_proj','gate_proj'): setattr(lm.model.language_model,name,nn.Linear(2,2)) setattr(lm.model.visual,name,nn.Linear(2,2)) targets=lora_targets(lm) self.assertEqual(len(targets),7) self.assertTrue(all(x.startswith('model.language_model.') for x in targets)) def test_invalid_target(self): r=self.row();r['target']=[.5,.9] with self.assertRaises(ValueError):normalize_row(r) if __name__=='__main__':unittest.main()