opensysone / source /tests /test_training_harness.py
andyshu's picture
Back up verified OpenSysOne training snapshot and pinned source
1a0a7fb verified
Raw History Blame
23.8 kB
import json
from argparse import Namespace
from pathlib import Path
import shutil
import tempfile
import threading
from types import SimpleNamespace
import unittest
from unittest.mock import patch
import urllib.error
import urllib.request
import torch
from torch import nn
import experiment
from experiment import (metrics, bootstrap_difference, decision_backward, restore_warm_start,
validation_fits, TrainingValidationInterrupted)
from jev_harness import compile_request, create_server, format_response, RemoteBackend, validate_response
from training_model import ADAPTER_VERSION, PROMPT_VERSION, LowRankLinear, TrainableScorer
class TrainingHarnessTests(unittest.TestCase):
def test_warm_start_changes_optimization_but_checks_weights_and_provenance(self):
scorer = TrainableScorer.__new__(TrainableScorer)
nn.Module.__init__(scorer)
scorer.lm = LowRankLinear(nn.Linear(7,5),rank=3,alpha=6)
scorer.head = nn.Linear(5,1)
scorer.provenance = {'revision':'pinned-base'}
config = {'rank':3,'alpha':6,'adapters':True,'lr':3e-5,'schedule_steps':7500,'seed':432}
saved = {'format':'opensysone-adapter-v1','prompt_version':PROMPT_VERSION,
'adapter_version':ADAPTER_VERSION,'model_provenance':scorer.provenance,
'data_signature':'frozen-data','config':{**config,'lr':1e-4,'schedule_steps':None,'seed':431},
'trainable_state':scorer.trainable_state(),
'optimizer':{'deliberately':'not a compatible optimizer state'}}
with torch.no_grad():
for parameter in scorer.parameters():
if parameter.requires_grad:
parameter.add_(1)
restore_warm_start(scorer,saved,config,'frozen-data')
for name,value in scorer.trainable_state().items():
self.assertTrue(torch.equal(value,saved['trainable_state'][name]))
for changed in ({'data_signature':'wrong-data'}, {'model_provenance':{'revision':'other-base'}},
{'prompt_version':'other-prompt'}, {'adapter_version':'other-adapter'},
{'config':{**saved['config'],'rank':4}}):
with self.assertRaises(ValueError):
restore_warm_start(scorer,{**saved,**changed},config,'frozen-data')
def test_expansion_opt_in_checks_lineage_before_restoring_weights(self):
scorer = SimpleNamespace(provenance={'revision':'pinned-base'})
from unittest.mock import Mock
scorer.restore_trainable = Mock()
config = {'rank':3,'alpha':6,'adapters':True,'allow_train_data_change':True}
saved = {'format':'opensysone-adapter-v1','prompt_version':PROMPT_VERSION,
'adapter_version':ADAPTER_VERSION,'model_provenance':scorer.provenance,
'data_signature':'original','config':config,'trainable_state':{'trusted':'weights'}}
with patch.object(experiment, 'verify_train_data_transition', side_effect=ValueError('reserved data changed')):
with self.assertRaisesRegex(ValueError, 'reserved data changed'):
restore_warm_start(scorer, saved, config, 'expanded')
scorer.restore_trainable.assert_not_called()
proof = {'kind':'training_split_only','parent_data_signature':'original','data_signature':'expanded'}
with patch.object(experiment, 'verify_train_data_transition', return_value=proof) as guard:
self.assertEqual(restore_warm_start(scorer, saved, config, 'expanded'), proof)
guard.assert_called_once_with(saved, config, 'expanded')
scorer.restore_trainable.assert_called_once_with(saved['trainable_state'])
for arguments in (Namespace(allow_train_data_change=True, warm_start=None, resume='checkpoint.pt'),
Namespace(allow_train_data_change=True, warm_start=None, resume=None)):
with self.assertRaisesRegex(ValueError, 'requires --warm-start'):
experiment.train(arguments)
def test_warm_started_training_uses_fresh_optimizer_and_publishes_evidence_first(self):
class TinyScorer(nn.Module):
def __init__(self,*args,**kwargs):
super().__init__()
# The pinned pretrained base is independent of the experiment
# seed; only fresh adapters/head initialization uses that seed.
with torch.random.fork_rng(devices=[]):
torch.manual_seed(777)
base = nn.Linear(7,5)
self.lm = LowRankLinear(base,rank=3,alpha=6)
self.lm.config = SimpleNamespace(attention_dropout=0)
self.head = nn.Linear(5,1)
self.provenance = {'revision':'pinned-base'}
self.adapter_names = ['lm']
self.branch_batch_size = 1
self.device = torch.device('cpu')
def score_examples(self,rows):
return [self.head(self.lm(torch.tensor(row['_sequences'],dtype=torch.float32))).flatten() for row in rows]
trainable_state = TrainableScorer.trainable_state
restore_trainable = TrainableScorer.restore_trainable
torch.manual_seed(431)
parent = TinyScorer()
parent_state = parent.trainable_state()
rows = [{'id':str(i),'group':str(i),'family':'tiny','choices':['a','b'],'target':i%2,
'_sequences':torch.randn(2,7).tolist()} for i in range(4)]
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
parent_file = root/'parent.pt'
torch.save({'format':'opensysone-adapter-v1','prompt_version':PROMPT_VERSION,
'adapter_version':ADAPTER_VERSION,'model_provenance':parent.provenance,
'data_signature':'frozen-data','config':{'rank':3,'alpha':6,'adapters':True},
'trainable_state':parent_state,'step':40,'source_commit':'parent-source',
'optimizer':{'deliberately':'unloadable'},'random_state':'must not be loaded'},parent_file)
args = Namespace(command='train',model=str(root/'model'),dataset=str(root/'data'),
output=str(root/'run'),resume=None,warm_start=str(parent_file),head_only=False,
seed=432,rank=3,alpha=6,max_tokens=32,branch_batch_size=1,two_pass=True,
lr=3e-5,head_lr=3e-5,epochs=3,effective_batch=2,steps=1,schedule_steps=7500,
validation_per_family=4,deadline=None,save_steps=250,save_seconds=900,eval_steps=1,patience=8)
saved_checkpoints = []
actual_save = experiment.save_torch
def checked_save(path,value):
if path.name == 'checkpoint.pt':
saved_checkpoints.append({'step':value['step'],'optimizer_steps':[
float(state['step']) for state in value['optimizer']['state'].values()]})
if path.name == 'best.pt':
evidence = [path.parent/f"validation_step_{value['step']:06d}_predictions.json"]
if value['step']==0:
evidence.append(path.parent/'initial_validation_predictions.json')
self.assertTrue(any(p.exists() for p in evidence),
'Best checkpoint was published before prediction evidence')
actual_save(path,value)
with (patch.object(experiment,'TrainableScorer',TinyScorer),
patch.object(experiment,'guard_memory'),patch.object(experiment.signal,'signal'),
patch.object(experiment,'STOP',False),
patch.object(experiment,'data_for',return_value=({'train':rows,'validation':rows},'frozen-data')),
patch.object(experiment,'save_torch',side_effect=checked_save),
patch.object(torch.cuda,'get_device_name',return_value='CPU test'),
patch.object(torch.cuda,'get_device_capability',return_value=(0,0)),
patch.object(torch.cuda,'get_rng_state_all',return_value=[]),
patch.object(torch.cuda,'max_memory_allocated',return_value=0),
patch.object(torch.cuda,'max_memory_reserved',return_value=0)):
experiment.train(args)
self.assertEqual(saved_checkpoints[0],{'step':0,'optimizer_steps':[]})
self.assertTrue(all(step==1 for step in saved_checkpoints[-1]['optimizer_steps']))
manifest = json.loads((root/'run/manifest.json').read_text())
self.assertEqual(manifest['initialization']['parent_step'],40)
self.assertEqual(manifest['initialization']['parent_checkpoint_sha256'],experiment.sha256(parent_file))
self.assertFalse(manifest['initialization']['restores_optimizer'])
self.assertFalse(manifest['initialization']['restores_rng'])
checkpoint = torch.load(root/'run/checkpoint.pt',weights_only=False)
self.assertEqual(checkpoint['config']['schedule_steps'],7500)
self.assertEqual(checkpoint['config']['lr'],3e-5)
self.assertEqual(checkpoint['step'],1)
initial = json.loads((root/'run/initial_validation_predictions.json').read_text())
parent.eval()
for expected,actual in zip(parent.score_examples(rows),initial):
self.assertEqual(expected.detach().tolist(),actual['logits'])
# A legacy intermediary retained the best checkpoint but omitted its
# prediction JSON. Resume must recover that evidence via provenance.
legacy = root/'legacy'
legacy.mkdir()
for filename in ('checkpoint.pt','best.pt'):
shutil.copy2(root/'run'/filename,legacy/filename)
# Make the resumable current step better than its inherited best.
# A uniform scorer beats this fixture's original overconfident head.
changed = torch.load(legacy/'checkpoint.pt',weights_only=False)
changed['trainable_state']['head.weight'].zero_()
changed['trainable_state']['head.bias'].zero_()
torch.save(changed,legacy/'checkpoint.pt')
resumed_args = Namespace(**{**vars(args),'output':str(root/'resumed'),
'resume':str(legacy/'checkpoint.pt'),'warm_start':None,'steps':2})
with (patch.object(experiment,'TrainableScorer',TinyScorer),
patch.object(experiment,'guard_memory'),patch.object(experiment.signal,'signal'),
patch.object(experiment,'STOP',False),
patch.object(experiment,'data_for',return_value=({'train':rows,'validation':rows},'frozen-data')),
patch.object(experiment,'save_torch',side_effect=checked_save),
patch.object(torch.cuda,'get_device_name',return_value='CPU test'),
patch.object(torch.cuda,'get_device_capability',return_value=(0,0)),
patch.object(torch.cuda,'get_rng_state_all',return_value=[]),
patch.object(torch.cuda,'set_rng_state_all'),
patch.object(torch.cuda,'max_memory_allocated',return_value=0),
patch.object(torch.cuda,'max_memory_reserved',return_value=0)):
experiment.train(resumed_args)
original_step = torch.optim.AdamW.step
def stopping_step(optimizer,*step_args,**step_kwargs):
result = original_step(optimizer,*step_args,**step_kwargs)
experiment.request_stop()
return result
stopped_args = Namespace(**{**vars(resumed_args),'output':str(root/'stopped'),
'resume':str(root/'resumed/checkpoint.pt'),'steps':3})
with patch.object(torch.optim.AdamW,'step',new=stopping_step):
experiment.train(stopped_args)
# Re-score real retained artifacts under the new validation-only
# policy without resetting Adam/RNG or changing trained weights.
experiment.STOP = False
crossfit_args = Namespace(**{**vars(stopped_args),'output':str(root/'crossfit'),
'resume':str(root/'stopped/checkpoint.pt'),
'selection_metric':experiment.SELECTION_METRIC})
experiment.train(crossfit_args)
resumed = torch.load(root/'resumed/checkpoint.pt',weights_only=False)
self.assertEqual(resumed['step'],2)
self.assertTrue(resumed['initialization']['restores_optimizer'])
self.assertTrue(resumed['initialization']['restores_rng'])
self.assertTrue(all(float(s['step'])==2 for s in resumed['optimizer']['state'].values()))
self.assertEqual(json.loads((root/'resumed/validation_step_000000_predictions.json').read_text()),initial)
best = torch.load(root/'resumed/best.pt',weights_only=False)
self.assertEqual(best['step'],1)
self.assertAlmostEqual(best['best_validation_macro_nll'],0.69314718056,places=6)
self.assertTrue((root/'resumed/validation_step_000001_predictions.json').exists())
stopped = json.loads((root/'stopped/summary.json').read_text())
self.assertEqual(stopped['status'],'interrupted')
self.assertEqual(stopped['final_correctness_status'],'skipped_on_stop')
self.assertIsNone(stopped['correctness'])
self.assertFalse((root/'stopped/correctness_final.json').exists())
self.assertEqual(torch.load(root/'stopped/checkpoint.pt',weights_only=False)['step'],3)
crossfit = torch.load(root/'crossfit/checkpoint.pt',weights_only=False)
stopped_checkpoint = torch.load(root/'stopped/checkpoint.pt',weights_only=False)
self.assertEqual(crossfit['selection_metric'],experiment.SELECTION_METRIC)
self.assertEqual(crossfit['initialization']['selection_policy_change']['from'],'raw_nll')
self.assertEqual(crossfit['step'],3)
for key,value in crossfit['trainable_state'].items():
self.assertTrue(torch.equal(value,stopped_checkpoint['trainable_state'][key]))
self.assertTrue(all(float(s['step'])==3 for s in crossfit['optimizer']['state'].values()))
selected = torch.load(root/'crossfit/best.pt',weights_only=False)
evidence = json.loads((root/'crossfit/best_validation_predictions.json').read_text())
self.assertAlmostEqual(selected['best_validation_selection_score'],
experiment.validation_selection(evidence)['score'],places=12)
def test_final_correctness_yields_between_forwards_and_restores_chunk_size(self):
class Scorer:
branch_batch_size = 3
calls = 0
def eval(self): pass
def score_examples(self,rows):
self.calls += 1
if self.calls == 2:
experiment.request_stop()
return [torch.tensor([1.0,2.0]) for row in rows]
scorer = Scorer()
rows = [{'family':'f','choices':['a','b'],'_sequences':[[1],[2]]}]
with patch.object(experiment,'STOP',False):
with self.assertRaisesRegex(TrainingValidationInterrupted,'correctness stopped'):
experiment.correctness(scorer,rows,allow_stop=True)
self.assertEqual(scorer.calls,2)
self.assertEqual(scorer.branch_batch_size,3)
def test_training_validation_yields_at_deadline_without_affecting_finalization(self):
class Scorer:
def eval(self): pass
def score_examples(self,rows): return [torch.tensor([1.0,2.0]) for row in rows]
rows = [{'id':'one','group':'g','family':'f','target':1,'choices':['a','b']}]
self.assertTrue(validation_fits(1000,300,now=500))
self.assertFalse(validation_fits(1000,300,now=600))
with patch.object(experiment,'STOP',False),patch.object(experiment.time,'time',return_value=1000):
with self.assertRaises(TrainingValidationInterrupted):
experiment.predict(Scorer(),rows,deadline=1000)
with patch.object(experiment,'STOP',True):
with self.assertRaises(TrainingValidationInterrupted):
experiment.predict(Scorer(),rows,deadline=float('inf'))
self.assertEqual(len(experiment.predict(Scorer(),rows)),1)
def test_http_timeout_and_retry_bounds(self):
for kwargs in ({'timeout':0},{'timeout':float('inf')},{'attempts':0},{'attempts':6}):
with self.assertRaises(ValueError):
RemoteBackend(api_key='test-only',**kwargs)
def test_two_pass_matches_categorical_gradients(self):
class TinyScorer(nn.Module):
def __init__(self):
super().__init__()
self.layer=LowRankLinear(nn.Linear(7,5),rank=3)
self.head=nn.Linear(5,1)
def score_examples(self,rows):
return [self.head(self.layer(torch.tensor(row['_sequences'],dtype=torch.float32))).flatten() for row in rows]
torch.manual_seed(32)
scorer=TinyScorer()
row={'_sequences':torch.randn(4,7).tolist(),'target':2}
ordinary_loss=decision_backward(scorer,row,divisor=3)
gradients={name:p.grad.clone() for name,p in scorer.named_parameters() if p.requires_grad}
scorer.zero_grad(set_to_none=True)
recomputed_loss=decision_backward(scorer,row,divisor=3,two_pass=True)
self.assertAlmostEqual(ordinary_loss,recomputed_loss,places=6)
for name,p in scorer.named_parameters():
if p.requires_grad:
torch.testing.assert_close(p.grad,gradients[name],atol=1e-6,rtol=1e-5)
def test_chat_template_returns_actual_token_ids(self):
from tokenizers import Tokenizer, models, pre_tokenizers
from transformers import PreTrainedTokenizerFast
backend = Tokenizer(models.WordLevel({'[UNK]':0,'yes':1,'no':2,'ASSISTANT':3},unk_token='[UNK]'))
backend.pre_tokenizer = pre_tokenizers.Whitespace()
tokenizer = PreTrainedTokenizerFast(tokenizer_object=backend,unk_token='[UNK]')
tokenizer.chat_template = "{{ messages[0]['content'] }} {{ messages[1]['content'] }} ASSISTANT"
scorer = TrainableScorer.__new__(TrainableScorer)
nn.Module.__init__(scorer)
scorer.tokenizer,scorer.max_tokens = tokenizer,768
row={'state':'Evidence is present','question':'Is the evidence present?','choices':['yes','no']}
sequences=scorer.sequences(row)
self.assertEqual(len(sequences),2)
self.assertGreater(len(sequences[0]),20)
self.assertTrue(all(isinstance(t,int) for s in sequences for t in s))
self.assertNotEqual(sequences[0],sequences[1])
scorer.max_tokens=4
with self.assertRaisesRegex(ValueError,'no truncation'):
scorer.sequences(row)
def test_adapter_preserves_base_and_learns(self):
torch.manual_seed(12)
base = nn.Linear(7, 5)
layer = LowRankLinear(base, rank=3, alpha=6)
values = torch.randn(2, 4, 7)
original = base(values).detach().clone()
self.assertTrue(torch.equal(original, layer(values)))
optimizer = torch.optim.SGD([p for p in layer.parameters() if p.requires_grad], lr=0.05)
loss = layer(values).square().mean()
loss.backward()
self.assertIsNone(base.weight.grad)
self.assertGreater(layer.adapter_b.grad.abs().sum().item(), 0)
optimizer.step()
self.assertFalse(torch.equal(original, layer(values)))
self.assertTrue(torch.equal(original, base(values)))
def test_checkpoint_restores_every_trainable_tensor(self):
scorer = TrainableScorer.__new__(TrainableScorer)
nn.Module.__init__(scorer)
scorer.lm = LowRankLinear(nn.Linear(7, 5), rank=3)
scorer.head = nn.Linear(5, 1)
state = scorer.trainable_state()
with torch.no_grad():
for p in scorer.parameters():
if p.requires_grad:
p.add_(1)
scorer.restore_trainable(state)
for name,p in scorer.named_parameters():
if p.requires_grad:
self.assertTrue(torch.equal(p,state[name]))
incomplete = dict(state)
incomplete.pop(next(iter(incomplete)))
with self.assertRaises(ValueError):
scorer.restore_trainable(incomplete)
def test_jev_shapes_and_score_expectation(self):
payload = json.loads(Path("examples/jev_request.json").read_text())
compiled = compile_request(payload)
response = format_response(compiled, [[0.2,0.8],[0.7,0.2,0.1],[0.1,0.3,0.6]])
validate_response(payload,response)
self.assertEqual(response['answers']['urgency']['noul'],0.8)
self.assertEqual(response['answers']['department']['choice'],'technical')
self.assertAlmostEqual(response['answers']['frustration']['score'],1.5)
renamed = {**payload,'questions':{'changed':payload['questions']['urgency']}}
self.assertEqual(compile_request(renamed)[0][4],compiled[0][4])
with self.assertRaises(ValueError):
compile_request({**payload,'questions':{'bad':{'type':'score','instructions':'Rate','criteria':['same','same']}}})
def test_http_hosted_client_and_authentication(self):
payload = json.loads(Path('examples/jev_request.json').read_text())
def backend(value):
compiled = compile_request(value)
return format_response(compiled,[[1/len(item[2])]*len(item[2]) for item in compiled])
server = create_server(backend,port=0,api_key='unit-test-key')
thread = threading.Thread(target=server.serve_forever,daemon=True)
thread.start()
url = f'http://127.0.0.1:{server.server_address[1]}'
try:
response = RemoteBackend(url,api_key='unit-test-key')(payload)
validate_response(payload,response)
with self.assertRaises(RuntimeError):
RemoteBackend(url,api_key='wrong')(payload)
request = urllib.request.Request(url+'/v1/systemone',b'{"bad": true}',
{'Authorization':'Bearer unit-test-key','Content-Type':'application/json'})
with self.assertRaises(urllib.error.HTTPError) as error:
urllib.request.urlopen(request)
self.assertEqual(error.exception.code,422)
finally:
server.shutdown()
server.server_close()
thread.join()
def test_metrics_use_stable_logs_and_grouped_bootstrap(self):
row = {'id':'one','group':'group','family':'family','target':1,'choices':['a','b'],
'probabilities':[1.0,0.0],'log_probabilities':[0.0,-1000.0]}
self.assertEqual(metrics([row])['nll'],1000)
tuned = {**row,'probabilities':[0.0,1.0],'log_probabilities':[-1000.0,0.0]}
delta = bootstrap_difference([row],[tuned],repetitions=20)
self.assertEqual(delta['nll'],[-1000.0,-1000.0])
self.assertEqual(delta['accuracy'],[1.0,1.0])
def test_group_bootstrap_preserves_decision_weighted_estimate(self):
row = {'id':'one','group':'small','family':'family','target':1,'choices':['a','b'],
'probabilities':[1.0,0.0],'log_probabilities':[0.0,-1000.0]}
base = [row]+[{**row,'id':str(i),'group':'large'} for i in range(3)]
tuned = [{**row,'probabilities':[0.0,1.0],'log_probabilities':[-1000.0,0.0]}]+base[1:]
delta = bootstrap_difference(base,tuned,repetitions=20)
self.assertEqual(delta['point_delta']['accuracy'],0.25)
self.assertEqual(delta['point_delta']['nll'],-250.0)
if __name__ == '__main__':
unittest.main()