opensysone / source /tests /test_profile_inference.py
andyshu's picture
Organize verified OpenSysOne publication payload
294f8ea verified
Raw History Blame
13.5 kB
import copy
import json
from pathlib import Path
import tempfile
from types import SimpleNamespace
import unittest
from unittest.mock import Mock, patch
import torch
from scripts import profile_inference as profile
class BoundaryTokenizer:
"""A tiny chat tokenizer with controllable continuation-boundary failures."""
def __init__(self, failure=None):
self.failure = failure
def apply_chat_template(self, messages, tokenize, **unused):
rendered = '|'.join(item['content'] for item in messages) + '|assistant|'
return [1, 2, 3] if tokenize else rendered
def encode(self, value, **unused):
if value.endswith('|assistant|'):
return [1, 2, 9] if self.failure == 'prefix' else [1, 2, 3]
letter = value[-1]
identifier = 4 + ord(letter) - ord('A')
if self.failure == 'merge':
return [1, 2, identifier]
if self.failure == 'multiple':
return [1, 2, 3, identifier, 0]
if self.failure == 'collision':
identifier = 4
return [1, 2, 3, identifier]
def decision(index=0, family='arc', **kwargs):
return {'id': f'{family}:{index}', 'group': f'{family}:group:{index}',
'family': family, 'state': 'A red package is on the table.',
'question': 'Which object is red?', 'choices': ['the package', 'the table'],
'target': 0, **kwargs}
class ProfileInferenceTests(unittest.TestCase):
def test_label_proof_rejects_boundary_merges_and_id_collisions(self):
row = decision(choices=['first', 'second', 'third'])
proof = profile.checked_label_encoding(BoundaryTokenizer(), row, 1024)
self.assertEqual(proof['labels'], ['A', 'B', 'C'])
self.assertEqual(proof['label_token_ids'], [4, 5, 6])
self.assertTrue(proof['boundary_checked'])
for failure in ('prefix', 'merge', 'multiple', 'collision'):
with self.subTest(failure=failure), self.assertRaises(ValueError):
profile.checked_label_encoding(BoundaryTokenizer(failure), row, 1024)
with self.assertRaisesRegex(ValueError, 'no truncation'):
profile.checked_label_encoding(BoundaryTokenizer(), row, 2)
def test_prompt_maps_labels_to_exact_candidate_order(self):
row = decision(state='literal\nSTATE value', question='literal QUESTION?',
choices=['Unicode café', 'multi\nline candidate', 'last'])
labels, messages = profile.label_prompt(row)
self.assertEqual(labels, list('ABC'))
self.assertIn('STATE:\nliteral\nSTATE value', messages[1]['content'])
self.assertIn('QUESTION:\nliteral QUESTION?', messages[1]['content'])
self.assertIn('A. Unicode café\nB. multi\nline candidate\nC. last', messages[1]['content'])
with self.assertRaises(ValueError):
profile.label_prompt(decision(choices=['only one']))
def test_indexed_label_projection_equals_full_vocabulary_conditioning(self):
torch.manual_seed(812)
hidden = torch.randn(1, 3, 7)
head = torch.nn.Linear(7, 23, bias=True)
tokenizer = BoundaryTokenizer()
row = decision(choices=['candidate B text', 'candidate A text', 'candidate C'])
proof = profile.checked_label_encoding(tokenizer, row, 1024)
observed = []
def forward(**kwargs):
observed.append(kwargs)
return SimpleNamespace(last_hidden_state=hidden)
scorer = SimpleNamespace(tokenizer=tokenizer, max_tokens=1024, device=torch.device('cpu'),
lm=SimpleNamespace(model=forward, lm_head=head))
result = profile.request_prediction(scorer, 'base_label', [row], [{'label': proof}])[0]
full_logits = head(hidden[:, -1, :])[0]
expected = torch.softmax(full_logits[proof['label_token_ids']], dim=0)
torch.testing.assert_close(torch.tensor(result['probabilities']), expected)
self.assertEqual(result['choices'], row['choices'])
self.assertEqual(result['predicted_choice'], row['choices'][int(expected.argmax())])
self.assertEqual(len(observed), 1)
self.assertFalse(observed[0]['use_cache'])
self.assertEqual(observed[0]['input_ids'].tolist(), [[1, 2, 3]])
mutated = {**row, 'state': 'a different state'}
with self.assertRaisesRegex(ValueError, 'boundary proof'):
profile.request_prediction(scorer, 'base_label', [mutated], [{'label': proof}])
def test_sample_is_balanced_group_unique_and_input_order_independent(self):
families = ('arc', 'banking', 'boolq', 'snli', 'social')
rows = [decision(index, family) for family in families for index in range(70)]
rows.extend([{**rows[0], 'id': 'duplicate-other-id'}, {**rows[1], 'id': 'duplicate-another'}])
selected = profile.balanced_sample(rows)
reversed_selected = profile.balanced_sample(list(reversed(rows)))
self.assertEqual([row['id'] for row in selected], [row['id'] for row in reversed_selected])
self.assertEqual(len({row['group'] for row in selected}), 320)
self.assertEqual(dict(profile.Counter(row['family'] for row in selected)), {family: 64 for family in families})
with self.assertRaisesRegex(ValueError, 'Insufficient'):
profile.balanced_sample([decision(0, family) for family in families])
def test_probability_and_timing_summaries_reject_invalid_measurements(self):
self.assertAlmostEqual(sum(profile.distribution([10000., 9999., -10000.])), 1.)
self.assertEqual(profile.distribution([0., 0.]), [.5, .5])
for logits in ([0., float('nan')], [0., float('inf')], [1.]):
with self.assertRaises(ValueError):
profile.distribution(logits)
summary = profile.timing_summary([float(value) for value in range(10, 0, -1)])
self.assertEqual(summary['median_seconds'], 5.5)
self.assertEqual(summary['p95_seconds'], 10.)
self.assertIn('exploratory', summary['p95_method'])
for samples in ([1.] * 9, [1.] * 9 + [float('nan')], [0.] * 10):
with self.assertRaises(ValueError):
profile.timing_summary(samples)
def test_repeat_gate_checks_distribution_and_choice_identity(self):
reference = [{'id': 'request', 'choices': ['a', 'b'], 'probabilities': [.7, .3]}]
self.assertEqual(profile.repeat_error(reference, copy.deepcopy(reference)), 0.)
for changes in ({'id': 'other'}, {'choices': ['b', 'a']}, {'probabilities': [.7]},
{'probabilities': [.699, .301]}, {'probabilities': [float('nan'), .3]}):
with self.subTest(changes=changes), self.assertRaises(ValueError):
profile.repeat_error(reference, [{**reference[0], **changes}])
def test_native_methods_retokenize_and_require_every_question_response(self):
calls = []
def score(rows):
calls.append(rows)
return [torch.tensor([3., 1.]) for _ in rows]
scorer = SimpleNamespace(score_examples=score, scores_token_baseline=score)
row = decision(_sequences=[[123]], _profile={'proof': True})
for method in ('trained', 'base_verifier'):
result = profile.request_prediction(scorer, method, [row], [{}])
self.assertEqual(result[0]['predicted_index'], 0)
self.assertNotIn('_sequences', calls[-1][0])
self.assertNotIn('_profile', calls[-1][0])
self.assertIn('_sequences', row)
with self.assertRaises(ValueError):
profile.request_prediction(scorer, 'trained', [row], [])
scorer.score_examples = lambda rows: []
with self.assertRaisesRegex(ValueError, 'question responses'):
profile.request_prediction(scorer, 'trained', [row], [{}])
def test_checkpoint_declaration_detects_changes_and_missing_expected_hash(self):
with tempfile.TemporaryDirectory() as temporary:
checkpoint = Path(temporary) / 'weights.pt'
checkpoint.write_bytes(b'frozen fixture weights')
expected = profile.sha(checkpoint)
profile.verify_checkpoint(checkpoint, expected)
with self.assertRaises(ValueError):
profile.verify_checkpoint(checkpoint, None)
checkpoint.write_bytes(b'changed fixture weights')
with self.assertRaisesRegex(ValueError, 'immutable'):
profile.verify_checkpoint(checkpoint, expected)
def test_supervisor_reaps_its_worker_after_state_write_failure(self):
with tempfile.TemporaryDirectory() as temporary:
root = Path(temporary)
(root / 'protocol.json').write_text('{}')
process = Mock(pid=123, returncode=-15)
process.poll.return_value = None
args = SimpleNamespace(prepared=str(root), output=str(root / 'output'), methods='trained',
checkpoint=None, only='accuracy', deadline='2099-01-01T00:00:00Z')
with patch.object(profile, 'verify_prepared', return_value=({}, {})), \
patch.object(profile.subprocess, 'Popen', return_value=process), \
patch.object(profile, 'json_write', side_effect=OSError('fixture disk failure')):
with self.assertRaisesRegex(OSError, 'disk failure'):
profile.run(args)
process.terminate.assert_called_once()
process.wait.assert_called_once_with(timeout=15)
process.kill.assert_not_called()
def test_protocol_is_frozen_before_any_heldout_row_read_and_checks_mutation(self):
from training_model import ADAPTER_VERSION, PROMPT_VERSION
with tempfile.TemporaryDirectory() as temporary:
root = Path(temporary)
model, data, output = root / 'model', root / 'data', root / 'prepared'
model.mkdir()
data.mkdir()
(model / 'tokenizer_config.json').write_text('{}')
splits = {'test': [decision(0, family) for family in ('arc', 'banking', 'boolq', 'snli')],
'holdout': [decision(0, 'social')], 'diagnostics': [decision(0, 'piqa')]}
for split, family in (('test', 'arc'), ('holdout', 'social'), ('diagnostics', 'piqa')):
splits[split].append(decision(1, family, state='too long for original eligibility'))
for name, rows in splits.items():
(data / (name + '.jsonl')).write_text(''.join(json.dumps(row) + '\n' for row in rows))
manifest = {'split_sha256': {name: profile.sha(data / (name + '.jsonl')) for name in ('test', 'holdout')},
'diagnostics': {'path': 'diagnostics.jsonl', 'sha256': profile.sha(data / 'diagnostics.jsonl')}}
(data / 'manifest.json').write_text(json.dumps(manifest))
artifact = {'format': 'opensysone-adapter-v1', 'adapter_version': ADAPTER_VERSION,
'prompt_version': PROMPT_VERSION, 'step': 3, 'model_provenance': {'revision': 'fixture'},
'config': {'model': str(model)}}
checkpoint = root / 'checkpoint.pt'
torch.save(artifact, checkpoint)
args = SimpleNamespace(checkpoint=str(checkpoint), dataset=str(data), output=str(output),
max_tokens=1024, accuracy_count=5, warmups=2, repeats=10, seed=917)
real_read = profile.read_rows
read_paths = []
def guarded_read(path):
frozen = json.loads((output / 'protocol.json').read_text())
self.assertEqual(frozen['accuracy_seed'], 917)
self.assertEqual(profile.sha(output / 'protocol.json'), (output / 'protocol.sha256').read_text().strip())
self.assertEqual(frozen['checkpoint_sha256'], profile.sha(checkpoint))
read_paths.append(Path(path).name)
return real_read(path)
scorer = SimpleNamespace(sequences=lambda row: [[1] * (513 if row['state'].startswith('too long') else 2)] * 2)
with patch.object(profile, 'read_rows', side_effect=guarded_read), \
patch.object(profile, 'tokenizer_only', return_value=scorer), \
patch.object(profile, 'synthetic_timing_cases', return_value=[]), \
patch.object(profile, 'compile_case', return_value={'fixture': True}):
profile.prepare(args)
self.assertEqual(read_paths, ['test.jsonl', 'holdout.jsonl', 'diagnostics.jsonl'])
protocol, requests = profile.verify_prepared(output)
self.assertEqual(len(requests['accuracy']['heldout']), 5)
self.assertEqual(len(requests['accuracy']['diagnostics']), 1)
preparation = json.loads((output / 'preparation.json').read_text())
self.assertEqual({row['split'] for row in preparation['filtered']}, {'test', 'holdout', 'diagnostics'})
self.assertTrue(all('512-token' in row['reason'] for row in preparation['filtered']))
(model / 'tokenizer_config.json').write_text('{"changed":true}')
with self.assertRaisesRegex(ValueError, 'configuration changed'):
profile.verify_prepared(output)
(model / 'tokenizer_config.json').write_text('{}')
(output / 'requests.json').write_text('{}')
with self.assertRaisesRegex(ValueError, 'requests changed'):
profile.verify_prepared(output)
if __name__ == '__main__':
unittest.main()