Merge nucleic/sleek-ember-seal-uady into dev

This commit is contained in:
2026-08-02 22:44:24 -07:00
parent ae616c6f85
commit 1a41febf73
4 changed files with 673 additions and 146 deletions
+55 -1
View File
@@ -5,6 +5,8 @@ from pathlib import Path
import mlx.core as mx
import mlx.nn as nn
import mlx.optimizers as optim
from mlx.utils import tree_flatten
MODULE_DIR = Path(__file__).resolve().parents[1]
@@ -18,7 +20,11 @@ from deep_model_mlx import (
save_weights,
)
from purpose_data import DataError, LABELS
from train_deep_mlx import _distillation_loss
from train_deep_mlx import (
_distillation_loss,
_load_optimizer_state,
_save_training_resume,
)
def tiny_config(*, checkpointing=False):
@@ -115,6 +121,54 @@ class DeepModelTests(unittest.TestCase):
self.assertAlmostEqual(0.0, float(value.item()), places=6)
self.assertEqual(teacher.shape, gradient.shape)
def test_training_resume_round_trips_model_and_optimizer(self):
class Tokenizer:
@staticmethod
def save_pretrained(destination):
(destination / "tokenizer_config.json").write_text("{}\n")
model = ModernBertForPurposeClassification(tiny_config())
optimizer = optim.AdamW(learning_rate=1e-3)
optimizer.init(model.trainable_parameters())
def loss(ids, mask):
return mx.mean(model(ids, mask)["purpose_logits"] ** 2)
value_and_grad = nn.value_and_grad(model, loss)
value, gradients = value_and_grad(
mx.array([[1, 3, 4, 2]]), mx.ones((1, 4), dtype=mx.int32)
)
optimizer.update(model, gradients)
mx.eval(value, model.parameters(), optimizer.state)
with tempfile.TemporaryDirectory() as temp:
checkpoint = _save_training_resume(
mx,
model,
optimizer,
Tokenizer(),
Path(temp),
{},
{"schemaVersion": 1, "status": "paused"},
)
restored_model = ModernBertForPurposeClassification(tiny_config())
restored_model.load_weights(
str(checkpoint / "model.safetensors"), strict=True
)
restored_optimizer = optim.AdamW(learning_rate=1e-3)
restored_optimizer.init(restored_model.trainable_parameters())
_load_optimizer_state(mx, restored_optimizer, checkpoint)
original = tree_flatten(optimizer.state, destination={})
restored = tree_flatten(restored_optimizer.state, destination={})
self.assertEqual(set(original), set(restored))
for key in original:
with self.subTest(optimizer_tensor=key):
self.assertEqual(
0,
float(mx.max(mx.abs(original[key] - restored[key])).item()),
)
def test_checkpoint_round_trip(self):
model = ModernBertForPurposeClassification(tiny_config())
with tempfile.TemporaryDirectory() as temp:
+77
View File
@@ -0,0 +1,77 @@
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()