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:
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user