78 lines
2.4 KiB
Python
78 lines
2.4 KiB
Python
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()
|