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