Download source/tests/test_profile_inference.py from andyshu/opensysone: direct link, hf CLI and curl.
- Browser
- Download file 13.5 kB
-
https://huggingface.co/andyshu/opensysone/resolve/294f8ea1b877ac86188aa88eade4b81f5f190293/source/tests/test_profile_inference.py
- Command line
-
hf download hf://andyshu/opensysone@294f8ea1b877ac86188aa88eade4b81f5f190293/source/tests/test_profile_inference.py
-
curl -L -o test_profile_inference.py https://huggingface.co/andyshu/opensysone/resolve/294f8ea1b877ac86188aa88eade4b81f5f190293/source/tests/test_profile_inference.py
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() | |