Repository navigation
fix(distill): flatten MFTLoss labels along with the logits - #2645
Conversation
MFTLoss.forward flattens both logit tensors to (batch * positions, vocab) and
documents that it only assumes the class dimension is last, so leading
dimensions are meant to be allowed. The labels were passed through untouched,
and _prepare_corrected_distributions then rejects anything that is not 1D:
ValueError: Logits must be a 2D tensor and labels must be a 1D tensor.
So the shapes a language model actually produces -- (batch, seq_len, vocab)
logits against (batch, seq_len) labels -- could not be used, which is the
setting Minifinetuning is for. A caller had to flatten the labels themselves,
and nothing said so.
Flatten them with the logits. The existing test drives MFTLoss through a vision
model, whose logits are already 2D and labels 1D, which is why this went unseen.
Signed-off-by: Siddhardha Nanda <99672439+SID-6921@users.noreply.github.com>
|
No actionable comments were generated in the recent review. 🎉 ℹ️ Recent review info⚙️ Run configuration
📒 Files selected for processing (1)
Included review availability: This review used your included allowance. Your plan provides up to 12 included reviews per hour; 11 remain after this review. 📝 WalkthroughWalkthrough
ChangesMFTLoss label handling
Priority: ⬇️ Low Estimated code review effort: 2 (Simple) | ~10 minutes Change: Bug fix Merge Risk: ⚪ Minimal · up to Sequence-shaped logits and labels are flattened in matching order before loss calculation, with a test comparing this path to explicitly flattened inputs. No concrete merge risk is established by the reviewed change. 🚥 Pre-merge checks | ✅ 5 | ❌ 1❌ Failed checks (1 warning)
✅ Passed checks (5 passed)
✨ Finishing Touches🧪 Generate unit tests (beta)
Comment |
|
@SID-6921 Thanks, but we can remove the changes to the |
AAnoosheh noted this fix isn't significant enough to warrant a changelog entry. Signed-off-by: Siddhardha Nanda <99672439+SID-6921@users.noreply.github.com>
|
Done in bef29e8 — changelog entry removed. |
|
The 4534 of 4535 tests passed. The one failure is Not pushing anything for it since it is out of scope here. Happy to rebase onto current |
|
@AAnoosheh thanks for the approval and the CHANGELOG note — addressed in bef29e8. The CI failure on this run isn't from this change: out of 4534 unit tests, the only failure is |
|
@AAnoosheh — thanks for the review and the vetting! Just flagging that the one failing check ( |
Codecov Report✅ All modified and coverable lines are covered by tests. Additional details and impacted files@@ Coverage Diff @@
## main #2645 +/- ##
==========================================
+ Coverage 72.16% 78.69% +6.53%
==========================================
Files 635 633 -2
Lines 70643 70639 -4
==========================================
+ Hits 50977 55590 +4613
+ Misses 19666 15049 -4617
Flags with carried forward coverage won't be shown. Click here to find out more. ☔ View full report in Codecov by Harness. 🚀 New features to boost your workflow:
|
|
Thanks for the approval, @AAnoosheh — this is green across the board now (linux, unit-pr-required-check, windows, the full multi-version matrix, code-quality, codecov). Whenever you have a moment to merge, it's ready. |
|
/ok to test bef29e8 |
|
Thanks for running The two red checks are one real failure plus its aggregator: Not pushing anything for it since it is out of scope here. |
|
This is green and approved aside from the known unrelated flake — ready to merge whenever convenient. |
|
/ok to test ce90633 |
What does this PR do?
Type of change: Bug fix
MFTLoss.forwardflattens both logit tensors to(batch * positions, vocab)and its docstring only assumes the class dimension is last, so leading dimensions are clearly meant to be allowed:The labels were passed through untouched, and
_prepare_corrected_distributionsrejects anything that is not 1-D:So the shapes a language model actually produces —
(batch, seq_len, vocab)logits against(batch, seq_len)labels — cannot be used:main(2, 8, 50)(2, 8)ValueError(16, 50)(16,)That is the setting Minifinetuning (arXiv:2506.15702) is for, so in practice a caller had to flatten the labels themselves and nothing documented that.
Flattening them alongside the logits fixes it, and the docstring now says what shape the labels are expected in.
Testing
test_mft_loss_accepts_sequence_shaped_logitsdrivesMFTLossat(2, 8, 50)/(2, 8)and asserts the result equals the pre-flattened(16, 50)/(16,)call. It fails onmainwith theValueErrorabove and passes here.Worth noting why this was not caught: the existing
test_distillation_model_mftdrivesMFTLossthrough a vision model, whose logits are already(batch, classes)and labels already(batch,), so the flattening never does anything there.tests/unit/torch/distillis 30 passed, 1 skipped locally. The skip isplugins/test_huggingface_kd.py, which needstransformersand is not installed in my environment; it exercisesLogitsDistillationLoss, which this change does not touch.Before your PR is "Ready for review"
CONTRIBUTING.md: N/AAdditional Information
Unrelated to my open ONNX PRs (#2553, #2554, #2567, #2575, #2583) — no shared files.
One thing I noticed next door and did not touch, in case it is of interest:
LogitsDistillationLossdefaults toreduction="mean", whileMFTLossin the same file defaults to"batchmean". PyTorch warns on every call that"mean"is not the KL divergence value and that it will be changed to behave as"batchmean"in a future major release, so the default path is currently a factor of the vocabulary size away from the other two reductions and will shift silently when that lands. Changing a training default is your call rather than mine, so I have left it alone — happy to open a separate issue if useful.Summary by CodeRabbit