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

This commit is contained in:
2026-07-31 14:30:10 -07:00
parent 4ed4763557
commit 3d35f4953f
5 changed files with 248 additions and 15 deletions
+2
View File
@@ -121,6 +121,8 @@ class TargetAndMetricTests(unittest.TestCase):
self.assertEqual(0, metrics["primaryHardSlice"]["accuracy"])
self.assertAlmostEqual(1 / 6, metrics["selectionScore"])
self.assertEqual(1, metrics["secondary"]["accuracy"])
self.assertEqual(["writing"], metrics["secondary"]["supportedLabels"])
self.assertEqual(1, metrics["secondary"]["supportedMacroRecall"])
self.assertEqual(1, metrics["mixed"]["f1"])
def test_mixed_threshold_prefers_higher_threshold_on_f1_tie(self):
+18
View File
@@ -18,6 +18,7 @@ from deep_model_mlx import (
save_weights,
)
from purpose_data import DataError, LABELS
from train_deep_mlx import _distillation_loss
def tiny_config(*, checkpointing=False):
@@ -97,6 +98,23 @@ class DeepModelTests(unittest.TestCase):
mx.eval(value, gradients)
self.assertTrue(float(value.item()) > 0)
def test_distillation_loss_matches_teacher_and_backpropagates(self):
teacher = mx.array([[2.0, 0.0, -1.0]])
def loss(student):
return _distillation_loss(
mx,
student,
teacher,
temperature=2.0,
weights=mx.ones((1,)),
)
value, gradient = mx.value_and_grad(loss)(teacher)
mx.eval(value, gradient)
self.assertAlmostEqual(0.0, float(value.item()), places=6)
self.assertEqual(teacher.shape, gradient.shape)
def test_checkpoint_round_trip(self):
model = ModernBertForPurposeClassification(tiny_config())
with tempfile.TemporaryDirectory() as temp: