Sitelet https://github.com/NVIDIA/Model-Optimizer/pull/2645
Skip to content

fix(distill): flatten MFTLoss labels along with the logits - #2645

Open
SID-6921 wants to merge 2 commits into
NVIDIA:mainfrom
SID-6921:fix/mft-loss-flatten-labels
Open

SID-6921 wants to merge 2 commits into
NVIDIA:mainfrom
SID-6921:fix/mft-loss-flatten-labels

Conversation

@SID-6921

@SID-6921 SID-6921 commented Oct 2, 2026 •

Copy link
Copy Markdown

What does this PR do?

Type of change: Bug fix

MFTLoss.forward flattens 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:

soft_log_probs = soft_log_probs.view(-1, soft_log_probs.size(-1))   # (new B, C)
target_logits  = target_logits.view(-1, target_logits.size(-1))     # (new B, C)
soft_targets = self._prepare_corrected_distributions(target_logits, labels, ...)

The labels were passed through untouched, and _prepare_corrected_distributions rejects anything that is not 1-D:

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 — cannot be used:

logits labels on main this PR
(2, 8, 50) (2, 8) ValueError converges, matches the flattened form exactly
(16, 50) (16,) works unchanged

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_logits drives MFTLoss at (2, 8, 50) / (2, 8) and asserts the result equals the pre-flattened (16, 50) / (16,) call. It fails on main with the ValueError above and passes here.

Worth noting why this was not caught: the existing test_distillation_model_mft drives MFTLoss through a vision model, whose logits are already (batch, classes) and labels already (batch,), so the flattening never does anything there.

tests/unit/torch/distill is 30 passed, 1 skipped locally. The skip is plugins/test_huggingface_kd.py, which needs transformers and is not installed in my environment; it exercises LogitsDistillationLoss, which this change does not touch.

Before your PR is "Ready for review"

  • Is this change backward compatible?: ✅ — already-1-D labels reshape to themselves, so existing callers are unaffected
  • If you copied code from any other sources or added a new PIP dependency, did you follow guidance in CONTRIBUTING.md: N/A
  • Did you write any new necessary tests?: ✅
  • Did you update Changelog?: ✅
  • Did you get Claude approval on this PR?: N/A

Additional 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: LogitsDistillationLoss defaults to reduction="mean", while MFTLoss in 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

  • Bug Fixes
    • Fixed MFTLoss to handle sequence-shaped labels alongside sequence-shaped logits, producing results consistent with equivalent flattened inputs.
  • Documentation
    • Clarified that labels should provide one value per logit position.

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>
Copilot AI balanced review requested due to automatic review settings October 2, 2026 20:36
@SID-6921
SID-6921 requested review from a team as code owners October 2, 2026 20:36
@copy-pr-bot

copy-pr-bot Bot commented Oct 2, 2026

Copy link
Copy Markdown

This pull request requires additional validation before any workflows can run on NVIDIA's runners.

Pull request vetters can view their responsibilities here.

Contributors can view more details about this message here.

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Copilot was unable to review this pull request because the user who requested the review has reached their quota limit.

@coderabbitai

coderabbitai Bot commented Oct 2, 2026 •

Copy link
Copy Markdown
Contributor

Review in Change Stack →

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 configuration

Configuration used: Repository: NVIDIA/Model-Optimizer/.coderabbit.yaml

Review profile: CHILL

Plan: Enterprise

Run ID: 21d43ad9-b0cf-4e9e-89f2-d0c8dbc2e4a1

📥 Commits

Reviewing files that changed from the base of the PR and between e4884ba and 108a1f5.

📒 Files selected for processing (3)
  • CHANGELOG.rst
  • modelopt/torch/distill/losses.py
  • tests/unit/torch/distill/test_distill.py

Included review availability: This review used your included allowance. Your plan provides up to 12 included reviews per hour; 11 remain after this review.


📝 Walkthrough

Walkthrough

MFTLoss 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.

Changes

MFTLoss label handling

Layer / File(s) Summary
Flatten labels and verify sequence-shaped inputs
modelopt/torch/distill/losses.py, tests/unit/torch/distill/test_distill.py, CHANGELOG.rst
MFTLoss flattens labels before preparing corrected distributions. Its documentation describes per-position labels. A test compares losses for sequence-shaped and flattened inputs, and the changelog records the change.

Priority: ⬇️ Low

Estimated code review effort: 2 (Simple) | ~8 minutes

Change: Bug fix

Merge Risk: ⚪ Minimal · up to 108a1

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)

Check name Status Explanation Resolution
Docstring Coverage ⚠️ Warning 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… Write docstrings for the functions missing them to satisfy the coverage threshold.
✅ Passed checks (5 passed)
Check name Status Explanation
Description Check ✅ Passed Check skipped - CodeRabbit’s high-level summary is enabled.
Title check ✅ Passed The title clearly and concisely describes the main change: flattening MFTLoss labels together with the logits.
Linked Issues check ✅ Passed Check skipped because no linked issues were found for this pull request.
Out of Scope Changes check ✅ Passed Check skipped because no linked issues were found for this pull request.
Security Anti-Patterns ✅ Passed PASS: The pull request changes only CHANGELOG.rst, modelopt/torch/distill/losses.py, and a distillation test. The added Python code only reshapes labels and adds test data. No added `torch.load(..…
Full details: Docstring Coverage

Explanation

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)
  • Create a new PR
  • Autopilot · Keep fixing CodeRabbit findings and required CI, and resolving merge conflicts

Autopilot is currently an internal CodeRabbit preview.


Comment @coderabbitai help to get the list of available commands.

@AAnoosheh

Copy link
Copy Markdown
Contributor

@SID-6921 Thanks, but we can remove the changes to the CHANGELOG.rst since it's not a significant change IMO

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>
@SID-6921

SID-6921 commented Oct 3, 2026

Copy link
Copy Markdown
Author

Done in bef29e8 — changelog entry removed.

@SID-6921

SID-6921 commented Oct 4, 2026

Copy link
Copy Markdown
Author

The linux / unit-pr-required-check failure is unrelated to this PR.

4534 of 4535 tests passed. The one failure is tests/unit/torch/speculative/plugins/test_hf_dflash2.py::TestDFlash2Forward::test_overfits_a_single_batch, which hit a pytest-timeout (>60s) — a speculative-decoding test with no connection to modelopt/torch/distill/losses.py. That test file does not exist on the main commit this branch was forked from, so it was added after, and the timeout looks like CI runner variance rather than anything this change touches.

Not pushing anything for it since it is out of scope here. Happy to rebase onto current main if that helps, if the test is still flaky there.

@SID-6921

SID-6921 commented Oct 5, 2026

Copy link
Copy Markdown
Author

@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 test_hf_dflash2.py::TestDFlash2Forward::test_overfits_a_single_batch, which timed out at 60s (pytest-timeout) — unrelated speculative-decoding code this PR doesn't touch. unit-pr-required-check just mirrors that job's exit code. Happy to re-run if useful, or let me know if you'd like anything else changed.

@SID-6921

SID-6921 commented Oct 5, 2026

Copy link
Copy Markdown
Author

@AAnoosheh — thanks for the review and the vetting! Just flagging that the one failing check (test_overfits_a_single_batch) is an unrelated flaky timeout, not from this diff — explained above. Since I can't re-run NVIDIA's CI myself, would you mind kicking off a re-run (or merging past it) whenever you get a chance?

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants