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