Skip to content

fix(megatron): preserve gradients across MTP depths - #1670

Closed
kkli08 wants to merge 1 commit into
areal-project:feat/megatron-mtp-cpfrom
kkli08:kkli08/mtp-multidepth-followup
Closed

kkli08 wants to merge 1 commit into
areal-project:feat/megatron-mtp-cpfrom
kkli08:kkli08/mtp-multidepth-followup

Conversation

@kkli08

@kkli08 kkli08 commented Sep 3, 2026

Copy link
Copy Markdown

Description

Stacked follow-up to #1659.

For multi-depth MTP, detaching hidden_states at every prediction depth prevents deeper auxiliary losses from updating earlier MTP blocks. This change:

  • detaches backbone hidden states only where they enter depth 0;
  • keeps later MTP depths connected, including repeated-layer mode;
  • continues detaching the shared embedding input at every depth;
  • preserves feat: add MTP training with speculative decoding rollout #1659's label/mask and tied/untied output-weight isolation;
  • adds exact-gradient tests for distinct and repeated MTP depths.

Related Issue

Related PR: #1659

Type of Change

  • 🐛 Bug fix
  • ✨ New feature
  • 💥 Breaking change
  • 📝 Documentation update
  • ♻️ Refactoring
  • ⚡ Performance improvement
  • ✅ Test coverage improvement

Checklist

  • I have read the Contributing Guide
  • Pre-commit hooks pass (pre-commit run --all-files)
  • Relevant tests pass; new tests added for new functionality
  • Documentation updated (not applicable; no user-facing API change)
  • Branch is up to date with main (this PR intentionally targets feat: add MTP training with speculative decoding rollout #1659's feature branch)
  • Self-reviewed via /review-pr command
  • This PR was created by a coding agent
  • This PR is a breaking change

Validation

  • 7 focused MTP CPU test cases passed with Megatron-Core 0.17.0 and Megatron-Bridge 0.4.0.
  • Exact gradients verified for distinct and repeated MTP depths.
  • An unclipped-SGD sanity check verifies that shared parameters receive no direct MTP gradient contribution while MTP parameters update.
  • Ruff format/check, Python compilation, and git diff --check passed.

The SGD check intentionally does not claim production optimizer-step equivalence: global gradient clipping can couple updates through the overall gradient norm. CUDA/Transformer Engine, PP/VPP, and distributed-optimizer execution remain for CI/GPU validation.

Signed-off-by: Ke Li <77185597+kkli08@users.noreply.github.com>
@kkli08

kkli08 commented Sep 3, 2026

Copy link
Copy Markdown
Author

Thanks @HT-Yuan and @sitabulaixizawaluduo. I’ve opened this focused follow-up as discussed. It adds correct gradient isolation for distinct and repeated multi-depth MTP on top of #1659, with exact-gradient tests. Broader Qwen3.5 GDN + MTP + CP support remains separate.

@sitabulaixizawaluduo

Copy link
Copy Markdown
Collaborator

Thanks @HT-Yuan and @sitabulaixizawaluduo. I’ve opened this focused follow-up as discussed. It adds correct gradient isolation for distinct and repeated multi-depth MTP on top of #1659, with exact-gradient tests. Broader Qwen3.5 GDN + MTP + CP support remains separate.

Can you provide a truly verified training curve to demonstrate the effectiveness of the current precision-recall?

@sitabulaixizawaluduo
sitabulaixizawaluduo deleted the branch areal-project:feat/megatron-mtp-cp September 6, 2026 03:39
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.

2 participants