Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user