OpenJudgement-4B-Preview / tests /test_runtime.py
xtristan's picture
Release OpenJudgement-4B-Preview
eb2384f
Raw History Blame Contribute Delete
5.7 kB
"""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()