diff --git a/modelopt/torch/distill/losses.py b/modelopt/torch/distill/losses.py index 832437b7e59..d094ac9e899 100644 --- a/modelopt/torch/distill/losses.py +++ b/modelopt/torch/distill/losses.py @@ -103,7 +103,9 @@ def forward( Args: logits_s: Student's logits, treated as prediction. logits_t: Teacher's logits, treated as training target. - labels: Labels for the ground truth, used to prepare the corrected teacher distributions. + labels: Labels for the ground truth, used to prepare the corrected teacher + distributions. Flattened alongside the logits, so it carries one label per + position, e.g. ``(batch, seq_len)`` against ``(batch, seq_len, vocab)`` logits. .. note:: @@ -115,7 +117,7 @@ def forward( target_logits: torch.Tensor = logits_t / self._temperature # (B, ..., C) target_logits = target_logits.view(-1, target_logits.size(-1)) # (new B, C) soft_targets = self._prepare_corrected_distributions( - target_logits, labels, self._threshold, apply_threshold_to_all=True + target_logits, labels.reshape(-1), self._threshold, apply_threshold_to_all=True ) kd_loss = F.kl_div( diff --git a/tests/unit/torch/distill/test_distill.py b/tests/unit/torch/distill/test_distill.py index 7c938aee90d..83bb179dfe0 100644 --- a/tests/unit/torch/distill/test_distill.py +++ b/tests/unit/torch/distill/test_distill.py @@ -127,6 +127,23 @@ def test_distillation_model_mft(): assert isinstance(loss, torch.Tensor) and loss.numel() == 1 +def test_mft_loss_accepts_sequence_shaped_logits(): + """MFTLoss flattens the logits it is given, so the labels have to follow them.""" + torch.manual_seed(0) + batch, seq_len, vocab = 2, 8, 50 + logits_s = torch.randn(batch, seq_len, vocab) + logits_t = torch.randn(batch, seq_len, vocab) + labels = torch.randint(0, vocab, (batch, seq_len)) + + loss = mtd.MFTLoss()(logits_s, logits_t, labels) + + # One label per position, so flattening first must not change the result. + flattened = mtd.MFTLoss()( + logits_s.reshape(-1, vocab), logits_t.reshape(-1, vocab), labels.reshape(-1) + ) + assert torch.allclose(loss, flattened) + + def test_distillation_mode_default_config(): student = tiny_mobilenet() with pytest.raises(AssertionError):