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: