import argparse from datetime import datetime, timedelta, timezone import hashlib import json import math from pathlib import Path import subprocess import tempfile import unittest from unittest.mock import patch import torch from scripts import final_validation as final from scripts.fleet_campaign import snapshot_candidate from selection import validation_selection, SELECTION_METRIC from training_model import ADAPTER_VERSION, PROMPT_VERSION class FakeScorer: device = torch.device('cpu') branch_batch_size = 1 def __init__(self, provenance, quality): self.provenance, self.quality = provenance, quality def sequences(self, row): return [[1, 2], [1, 3]] def eval(self): return self def score_examples(self, rows): return [torch.tensor([2.0 * self.quality, 0.0]) for row in rows] class FinalValidationTests(unittest.TestCase): def setUp(self): self.temp = tempfile.TemporaryDirectory() self.addCleanup(self.temp.cleanup) self.root = Path(self.temp.name) self.source, self.dataset, self.model = [self.root / name for name in ('training', 'dataset', 'model')] for path in (self.source, self.dataset, self.model): path.mkdir() self.rows = [{'id': f'{family}-{i}', 'group': f'{family}-{i//2}', 'family': family, 'choices': ['a', 'b'], 'target': 0, 'state': 'context', 'question': 'question'} for family in ('arc', 'banking', 'boolq', 'snli') for i in range(128)] predictions = [{**row, 'logits': [0.0, 0.0], 'probabilities': [0.5, 0.5], 'log_probabilities': [-math.log(2), -math.log(2)]} for row in self.rows] selection = validation_selection(predictions) for split in ('train', 'validation', 'calibration', 'test', 'holdout'): (self.dataset / (split + '.jsonl')).write_text( ''.join(json.dumps(row) + '\n' for row in self.rows) if split == 'validation' else 'opaque bytes, not JSON\n') split_hashes = {split: final.checksum(self.dataset / (split + '.jsonl')) for split in ('train', 'validation', 'calibration', 'test', 'holdout')} final.experiment.write_json(self.dataset / 'manifest.json', {'split_sha256': split_hashes}) self.provenance = {'model_id': 'fixture/scorer', 'revision': 'pin'} final.experiment.write_json(self.model / 'opensysone-provenance.json', self.provenance) config = {'dataset': str(self.dataset), 'model': str(self.model), 'output': str(self.source), 'rank': 8, 'alpha': 16, 'adapters': True, 'max_tokens': 512, 'branch_batch_size': 1, 'validation_per_family': 128} signature = hashlib.sha256(json.dumps({'data': split_hashes, 'model': self.provenance, 'implementation': final.checksum(final.ROOT / 'training_model.py'), 'max_tokens': 512}, sort_keys=True).encode()).hexdigest() revision = subprocess.check_output(['git', 'rev-parse', 'HEAD'], cwd=final.ROOT, text=True).strip() self.best = {'format': 'opensysone-adapter-v1', 'config': config, 'step': 0, 'prompt_version': PROMPT_VERSION, 'adapter_version': ADAPTER_VERSION, 'trainable_state': {'quality': torch.tensor(0.0)}, 'source_commit': revision, 'model_provenance': self.provenance, 'data_signature': signature, 'selection_metric': SELECTION_METRIC, 'best_validation_macro_nll': math.log(2), 'best_validation_selection_score': selection['score']} self.latest = {**self.best, 'step': 143, 'trainable_state': {'quality': torch.tensor(1.0)}, 'optimizer': {'must_not_restore': True}, 'random_state': (), 'torch_rng': torch.zeros(1), 'cuda_rng': []} torch.save(self.best, self.source / 'best.pt') torch.save(self.latest, self.source / 'checkpoint.pt') final.experiment.write_json(self.source / 'initial_validation_predictions.json', predictions) self.reference = self.root / 'reference.json' final.experiment.write_json(self.reference, predictions) self.manifest = {'pid': 999999999, 'git_commit': revision, 'config': config, 'data_signature': signature, 'model_provenance': self.provenance, 'source_sha256': {name: final.checksum(final.ROOT / name) for name in ('training_model.py', 'decision_model.py', 'selection.py')}} final.experiment.write_json(self.source / 'manifest.json', self.manifest) self.checks = {name + '_probability_max_abs': 0.0 for name in ('branch_chunks_1', 'branch_chunks_2', 'branch_chunks_4', 'question_isolation', 'candidate_permutation', 'repeat')} self.checks['tolerance_probability_abs'] = 1e-4 final.experiment.write_json(self.source / 'correctness_initial.json', self.checks) self.args = argparse.Namespace(checkpoint=str(self.source / 'checkpoint.pt'), reference_predictions=str(self.reference), reference_sha256=final.checksum(self.reference), evidence_dir=[], output=str(self.root / 'output'), deadline=(datetime.now(timezone.utc) + timedelta(minutes=5)).isoformat(), device='cpu') self.original_hashes = [final.checksum(self.source / name) for name in ('best.pt', 'checkpoint.pt')] def execute(self): def load(path, device): self.assertTrue(Path(path).is_file()) saved = torch.load(path, map_location='cpu', weights_only=False) self.assertTrue(all(key not in saved for key in final.STRIPPED_KEYS)) self.assertEqual(saved['source_commit'], self.latest['source_commit']) return FakeScorer(saved['model_provenance'], float(saved['trainable_state']['quality'])), saved with patch.object(final.experiment, 'guard_memory') as guard, \ patch.object(final.experiment, 'load_artifact', side_effect=load), \ patch.object(final.experiment, 'correctness', return_value=self.checks): result = final.run(self.args) guard.assert_called_once_with('cpu') self.assertEqual(self.original_hashes, [final.checksum(self.source / name) for name in ('best.pt', 'checkpoint.pt')]) return result def test_latest_is_promoted_without_updates_and_production_fleet_accepts_output(self): result = self.execute() self.assertTrue(result['promoted_latest']) self.assertEqual(result['selected_step'], 143) self.assertEqual(result['optimizer_updates'], 0) self.assertFalse(result['reserved_predictions_accessed']) out = Path(self.args.output) self.assertTrue((out / 'latest.pt').is_file()) self.assertTrue((out / 'latest.evaluated.pt').is_file()) chosen = torch.load(out / 'best.pt', map_location='cpu', weights_only=False) self.assertEqual(chosen['source_commit'], self.latest['source_commit']) self.assertEqual(chosen['final_validation']['original_checkpoint_sha256'], self.original_hashes[1]) record = snapshot_candidate({'name': 'fixture', 'host': 'local', 'campaign': str(self.root / 'stopped-campaign'), 'training': str(out)}, self.root / 'fleet-snapshot', final.read_json(self.reference)) self.assertTrue(record['eligible']) self.assertEqual(record['step'], 143) def test_no_improvement_preserves_original_best_bytes(self): self.latest['trainable_state']['quality'].zero_() torch.save(self.latest, self.source / 'checkpoint.pt') self.original_hashes[1] = final.checksum(self.source / 'checkpoint.pt') result = self.execute() self.assertFalse(result['promoted_latest']) self.assertEqual(result['best_sha256'], self.original_hashes[0]) self.assertEqual(result['selected_step'], 0) def test_threshold_is_strict_and_nonfinite_scores_rejected(self): self.assertFalse(final.choose_latest(1.0, 1.0 - 0.001)) self.assertTrue(final.choose_latest(1.0, 1.0 - 0.00101)) for value in (math.nan, math.inf, -math.inf): with self.assertRaises(ValueError): final.choose_latest(1.0, value) def test_changed_reference_fails_before_model_loading(self): self.args.reference_sha256 = '0' * 64 with patch.object(final.experiment, 'load_artifact') as loader: with self.assertRaisesRegex(ValueError, 'reference checksum'): final.run(self.args) loader.assert_not_called() self.assertEqual(final.read_json(Path(self.args.output) / 'state.json')['status'], 'failed') def test_missing_validation_row_rejected_instead_of_dropped(self): (self.dataset / 'validation.jsonl').write_text(''.join(json.dumps(row) + '\n' for row in self.rows[:-1])) with self.assertRaisesRegex(ValueError, 'exactly the frozen 512'): final.validation_rows(self.dataset, final.read_json(self.reference), FakeScorer(self.provenance, 1)) def test_inference_failure_preserves_durable_inputs_and_no_selected_output(self): with patch.object(final.experiment, 'guard_memory'), patch.object(final.experiment, 'load_artifact', side_effect=RuntimeError('fixture failure')): with self.assertRaisesRegex(RuntimeError, 'fixture failure'): final.run(self.args) out = Path(self.args.output) self.assertTrue((out / 'latest.pt').is_file()) self.assertTrue((out / 'source/checkpoint.pt').is_file()) self.assertFalse((out / 'best.pt').exists()) self.assertEqual(self.original_hashes, [final.checksum(self.source / name) for name in ('best.pt', 'checkpoint.pt')]) if __name__ == '__main__': unittest.main()