Skip to content

fix(mtp): preserve padding masks and handle empty auxiliary loss - #7611

Open
cuichenx wants to merge 4 commits into
NVIDIA:mainfrom
cuichenx:chcui/maya/mtp-padding-mask-roll
Open

cuichenx wants to merge 4 commits into
NVIDIA:mainfrom
cuichenx:chcui/maya/mtp-padding-mask-roll

Conversation

@cuichenx

@cuichenx cuichenx commented Sep 23, 2026 •

Copy link
Copy Markdown
Contributor
  • I, the PR author, have personally reviewed every line of this PR.

What does this PR do?

Keep padding and artificial MTP sequence-end positions out of MoE routing. Padding masks use True for padding, but the rolling helper fills newly exposed positions with zero. Rolling validity (~padding_mask) and inverting the result keeps these positions masked at every MTP depth.

Under sequence parallelism, gather the sharded mask over the explicit TP group before rolling in the full CP-local layout, then scatter it back. This preserves actual sequence boundaries instead of treating each TP shard as a separate sequence.

HybridModel now forwards its padding mask through MTP preparation and execution. When preparation changes the CP layout, convert mask validity alongside activations so newly inserted THD padding remains excluded. The precomputed-embedding path also shifts mask validity at each depth; it previously bypassed the token-embedding path's mask handling. Existing restrictions on padding masks in hybrid chunkwise CP remain unchanged.

Correct masking exposes zero-token auxiliary-loss normalization for entirely padded input. Clamp only the normalization denominator to at least one in fused and unfused paths, yielding zero loss/gradient without changing token counters or positive-count normalization. Replay failure labels now identify both fused and all_padding.

The token-embedding path adds one boolean-mask TP all-gather per masked MTP depth under SP, without an additional CP exchange relative to its original mask roll. The newly supported precomputed-embedding path adds a mask roll with the required TP/CP communication; Hybrid CP layout conversion also redistributes the supplied mask. Mask-free paths add no mask communication. No throughput claim.

Related: #7278 addresses logical-length truncation in packed CP rolling.

Validation

  • Previous sequence-parallel fix: six new TP2 cases failed before the fix; 30 focused MTP/rolling tests passed on each of four GB200 ranks in unchanged NeMo 26.10.rc0.
  • Hybrid follow-up in unchanged NeMo 26.10.rc0 on four GB200 GPUs: the previous head fails all eight masked Hybrid cases while eight mask-free controls pass. The fix passes 62 focused MTP tests on every rank, including all TP2/CP2 cases, with no skips.
  • Auxiliary-loss replay checks on every rank: 2 passed, 1 expected failure, 1 XPASS. The existing non-strict fused-kernel determinism marker accounts for the latter two; the all-padding fused case passes. The first combined-suite launch stopped during determinism-environment setup; separate suite invocations with the supported pinned-config fallback completed successfully.
  • Real small HybridModel tests run two MoE MTP depths and backward, checking router masks and finite router weight gradients. They cover token and precomputed embeddings, absent/present masks, fully valid/right-padded/all-padding rows, and TP1/2 × CP1/2.
  • Actual distributed CP layout conversion tests cover unpacked and padded THD layouts, inserted padding, absent masks, and input immutability. Integration tests verify the prepared mask reaches Hybrid MTP.
  • The unchanged auxiliary-loss fix previously passed 14 targeted router cases covering standard/per-sequence/global auxiliary loss, fused/unfused execution, and retained routing history.
  • Changed-file repository autoformatter and focused pre-commit passed. Advisory mypy retains 11 existing diagnostics outside the added code. Required all-file pre-commit reports unrelated existing formatting/lint issues in a disposable copy.
  • Independent follow-up review found no remaining actionable issues after the precomputed-embedding path was corrected.

Native tests import this upstream source with the existing container dependencies. A private fixture bypasses only the unrelated dataset download. These are synthetic unit/model tests, not Qwen/GLM training convergence or full-model gradient-invariance validation.

Pre-checks

  • Added relevant unit tests.
  • Added full-model functional tests (not part of this bounded mask fix).
  • Ran the repository autoformatter.
  • Independent review completed.

Signed-off-by: Chen Cui <chcui@nvidia.com>
@cuichenx cuichenx added Run MBridge tests Attach this for testing this PR against MBridge main bug fix labels Sep 23, 2026
@copy-pr-bot

copy-pr-bot Bot commented Sep 23, 2026

Copy link
Copy Markdown

Auto-sync is disabled for draft pull requests in this repository. Workflows must be run manually.

Contributors can view more details about this message here.

@cuichenx

Copy link
Copy Markdown
Contributor Author

/ok to test 88f8ff9

@ko3n1g

ko3n1g commented Sep 23, 2026

Copy link
Copy Markdown
Contributor

/review

@svcnvidia-nemo-ci svcnvidia-nemo-ci 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.

The padding-polarity fix has two blocking gaps in supported MTP/MoE paths: it misaligns masks under sequence parallelism, and HybridModel does not forward its padding mask into this code path.

Review generated by codex

Comment thread megatron/core/transformer/multi_token_prediction.py
Comment thread megatron/core/transformer/multi_token_prediction.py
@ko3n1g

ko3n1g commented Sep 23, 2026

Copy link
Copy Markdown
Contributor

/review model=claude

@svcnvidia-nemo-ci svcnvidia-nemo-ci 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.

Light review — fix(mtp): preserve padding masks and handle empty auxiliary loss

I verified both production changes against the head and base trees.

The polarity fix is the right diagnosis. padding_mask is True = padding everywhere on this path (GPTModel.forward docstring at megatron/core/models/gpt/gpt_model.py:619-621, Router.forward at megatron/core/transformer/moe/router.py:1065-1066, _apply_expert_bias's ~flat_mask, compute_routing_scores_for_aux_loss's valid_mask = (~padding_mask)), while every fill site in roll_tensor / _roll_tensor_packed_seq writes 0. Rolling validity and re-inverting therefore makes all three fill sites mean "padding": the CP-rank-0 mirrored-tail fill (multi_token_prediction.py:301), the last-rank splice (:315), and the packed document-boundary / physical-padding zeroing (:383-385). Base's direct roll marked those positions valid, admitting exactly the artificial terminal token the PR describes. Input tensors are not mutated (~ allocates, the packed path clones), and an absent mask stays absent.

The aux-loss clamp is behaviour-preserving where it matters. total_num_tokens only reaches 0 through the with_padding_mask=True branch of get_tokens_per_expert_and_token_count (moe_utils.py:290-292), where it is a 0-d tensor; the unfused path then divides by total_num_tokens**2, so 0 valid tokens produced 0 * inf = NaN. The mask-free path passes a Python int >= 1, so max(..., 1) is a no-op there and no golden values should move. Clamping rather than short-circuiting also keeps the reduce_from_tensor_model_parallel_region collective balanced across ranks, which is the right call. attach_and_log_load_balancing_loss only multiplies by valid_token_count, so no second division needs guarding.

One gap worth resolving before this lands for sequence-parallel users

GPTModel._preprocess scatters padding_mask across TP ranks when sequence_parallel=True (gpt_model.py:388-395), so the mask arriving at _get_embeddings is a [b, s/tp] shard while input_ids/position_ids are the full CP-local [b, s]. The roll consequently treats every shard as a standalone sequence, and with the new polarity the fill is True, so the last position of each SP shard is marked padding at every depth — up to depth genuinely valid tokens per shard per row drop out of the aux-loss / z-loss / expert-bias statistics. Base's zero-fill happened to be correct at those interior boundaries, so this is a change in the opposite direction for SP users rather than a strict improvement. Details and options inline.

The new MTP regression test is tp=1 only, and _create_config_and_mtp_block_spec enables sequence_parallel only when tp > 1, so nothing in the suite exercises this.

FYI, out of scope for this diff

Hybrid MTP is unaffected either way: HybridModel.forward passes padding_mask to the backbone but omits it from its self.mtp(...) call (megatron/core/models/hybrid/hybrid_model.py:745-759), so hybrid MTP+MoE layers never receive a mask at all. That is a pre-existing gap in the feature, not something this PR changes.

Test hygiene otherwise looks good: the independent length-based oracle (rather than re-deriving through roll_tensor) is the right shape for this test, the xfail inherited by the new all_padding cell is strict=False so an XPASS will not fail CI, and touching tests/unit_tests/determinism/kernels/test_moe_kernels.py satisfies tools/check_kernel_determinism_coverage.py for the moe_utils manifest entry.

Review generated by claude

Comment thread megatron/core/transformer/multi_token_prediction.py
Comment thread tests/unit_tests/transformer/test_multi_token_prediction.py Outdated
Comment thread tests/unit_tests/determinism/kernels/test_moe_kernels.py Outdated
Signed-off-by: Chen Cui <chcui@nvidia.com>
@cuichenx

Copy link
Copy Markdown
Contributor Author

/ok to test c7c4365

@cuichenx

Copy link
Copy Markdown
Contributor Author

/ok to test 993128b

@cuichenx cuichenx removed the Run MBridge tests Attach this for testing this PR against MBridge main label Sep 29, 2026
@cuichenx
cuichenx marked this pull request as ready for review September 29, 2026 23:29
@cuichenx
cuichenx requested review from a team as code owners September 29, 2026 23:29
@cuichenx

Copy link
Copy Markdown
Contributor Author

/review

@svcnvidia-nemo-ci svcnvidia-nemo-ci 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.

LGTM. No actionable issues found in this light review. Regression tests cover TP/CP mask rolling, hybrid MTP propagation, layout conversion, and all-padding auxiliary losses. Tests were not run in this read-only environment.

Review generated by codex

This branch was successfully deployed

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

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants