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>
|
Navigate logical layers of code changes, visualize relationships, and explore their blast radius. No actionable comments were generated in the recent review. 🎉 ℹ️ Recent review info⚙️ Run configurationConfiguration used: Repository: NVIDIA/Model-Optimizer/.coderabbit.yaml Review profile: CHILL Plan: Enterprise Run ID: 📒 Files selected for processing (3)
Included review availability: This review used your included allowance. Your plan provides up to 12 included reviews per hour; 11 remain after this review. 📝 WalkthroughWalkthroughMFTLoss now flattens labels before preparing corrected distributions. Documentation and a changelog entry describe the expected label shape. A test compares losses for sequence-shaped and flattened inputs. ChangesMFTLoss label handling
Priority: ⬇️ Low Estimated code review effort: 2 (Simple) | ~8 minutes Change: Bug fix Merge Risk: ⚪ Minimal · up to Sequence-shaped MFTLoss inputs now align labels with flattened logits, and the added unit test checks equivalence with flat inputs. No actionable merge-blocking risk is evident; proceed after normal checks. 🚥 Pre-merge checks | ✅ 5 | ❌ 1❌ Failed checks (1 warning)
✅ Passed checks (5 passed)
Full details: Docstring CoverageExplanation Docstring coverage is 50.00% which is insufficient. The required threshold is 80.00%. Docstring coverage is scoped to functions touched by this diff. Analyzed 4 functions across 2 files. (1 skipped: 1 unsupported.) ✨ 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 ( |
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
MFTLossto handle sequence-shaped labels alongside sequence-shaped logits, producing results consistent with equivalent flattened inputs.