Files
nucleic-purpose-classifier/tests/test_train_deep_resume.py
T

78 lines
2.4 KiB
Python
Raw Normal View History

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