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

This commit is contained in:
2026-07-31 03:05:29 -07:00
parent 9a1228efbb
commit 4ed4763557
4 changed files with 131 additions and 10 deletions
+35
View File
@@ -13,6 +13,7 @@ sys.path.insert(0, str(MODULE_DIR))
from deep_model_mlx import (
ModernBertForPurposeClassification,
ModernBertPurposeConfig,
load_checkpoint_weights,
load_pretrained_weights,
save_weights,
)
@@ -114,6 +115,40 @@ class DeepModelTests(unittest.TestCase):
0, float(mx.max(mx.abs(first[key] - second[key])).item())
)
def test_trained_checkpoint_loader_restores_every_task_head(self):
model = ModernBertForPurposeClassification(tiny_config())
with tempfile.TemporaryDirectory() as temp:
path = Path(temp) / "model.safetensors"
save_weights(model, path)
restored = ModernBertForPurposeClassification(tiny_config())
report = load_checkpoint_weights(restored, path)
self.assertEqual(0, report["freshTaskHeads"])
self.assertEqual(0, report["ignored"])
ids = mx.array([[1, 3, 4, 2]])
mask = mx.ones((1, 4), dtype=mx.int32)
first = model(ids, mask)
second = restored(ids, mask)
mx.eval(*first.values(), *second.values())
for key in first:
with self.subTest(head=key):
self.assertEqual(
0, float(mx.max(mx.abs(first[key] - second[key])).item())
)
def test_trained_checkpoint_loader_rejects_missing_task_heads(self):
model = ModernBertForPurposeClassification(tiny_config())
with tempfile.TemporaryDirectory() as temp:
complete = Path(temp) / "complete.safetensors"
partial = Path(temp) / "partial.safetensors"
save_weights(model, complete)
weights = mx.load(str(complete))
weights.pop("purpose_classifier.weight")
mx.save_safetensors(str(partial), weights)
with self.assertRaisesRegex(DataError, "missing 1 tensors"):
load_checkpoint_weights(
ModernBertForPurposeClassification(tiny_config()), partial
)
def test_pretrained_loader_rejects_partial_backbone(self):
with tempfile.TemporaryDirectory() as temp:
path = Path(temp) / "partial.safetensors"