import json import signal import sys import tempfile import unittest from pathlib import Path MODULE_DIR = Path(__file__).resolve().parents[1] sys.path.insert(0, str(MODULE_DIR)) from train_deep_mlx import ( RESUME_SCHEMA_VERSION, _ShutdownController, _read_resume_state, _restore_resume_arguments, build_parser, ) from purpose_data import DataError class DeepTrainingResumeTests(unittest.TestCase): def test_saved_arguments_make_resume_command_self_contained(self): parser = build_parser() args = parser.parse_args( [ "--resume-training", "/tmp/purpose-deep/resume", ] ) saved = { "variant": "base", "dataset_dir": "/datasets/v2", "epochs": 7, "distillation_cache": "/datasets/teacher.pt", "distillation_weight": 0.5, "progress_steps": 19, } checkpoint = Path("/tmp/purpose-deep/resume") _restore_resume_arguments( args, checkpoint, {"arguments": saved}, ) self.assertEqual(Path("/datasets/v2"), args.dataset_dir) self.assertEqual(Path("/datasets/teacher.pt"), args.distillation_cache) self.assertEqual(7, args.epochs) self.assertEqual(0.5, args.distillation_weight) self.assertEqual(checkpoint.parent, args.output_dir) self.assertEqual(checkpoint, args.resume_training) self.assertIsNone(args.resume_from) self.assertFalse(args.overwrite_output) def test_resume_state_fails_closed_on_wrong_schema(self): with tempfile.TemporaryDirectory() as temp: checkpoint = Path(temp) (checkpoint / "resume-state.json").write_text( json.dumps( { "schemaVersion": RESUME_SCHEMA_VERSION + 1, "status": "paused", } ) ) with self.assertRaisesRegex(DataError, "unsupported"): _read_resume_state(checkpoint) def test_second_shutdown_request_is_forceful(self): controller = _ShutdownController() controller._handle(signal.SIGINT, None) self.assertEqual(signal.SIGINT, controller.signum) with self.assertRaises(KeyboardInterrupt): controller._handle(signal.SIGINT, None) if __name__ == "__main__": unittest.main()