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