Skip to content

feat: add MTP training with speculative decoding rollout - #1659

Merged
sitabulaixizawaluduo merged 7 commits into
mainfrom
feat/megatron-mtp-cp
Sep 6, 2026
Merged

sitabulaixizawaluduo merged 7 commits into
mainfrom
feat/megatron-mtp-cp

Conversation

@HT-Yuan

@HT-Yuan HT-Yuan commented Sep 1, 2026 •

Copy link
Copy Markdown
Collaborator

Train the built-in MTP head jointly with GRPO/SFT on Megatron and use it during rollout through SGLang NEXTN speculative decoding. Online weight synchronization updates both the target model and the built-in MTP draft runner, ensuring that speculative decoding uses the latest RL-trained MTP weights instead of a stale initialization.

MTP is trained through Megatron-Core's auxiliary cross-entropy path with a configurable loss scaling factor of 0.1 by default. AReaL supplies independent MTP label and loss-mask channels while keeping the main forward path logits-based. Shared output weights, backbone hidden states, and embedding inputs are detached from the MTP loss graph, so the backbone receives only the policy/SFT gradient while the MTP-specific parameters learn from future-token supervision.

For packed THD training with context parallelism, AReaL applies the same per-sequence zigzag CP split and rank-local repacking to MTP labels and loss masks as it does to input IDs. Megatron-Core's CP-aware rolling then aligns future-token targets across CP ranks without crossing packed sequence boundaries.

For online rollout, a focused compatibility bridge for sglang==0.5.10.post1 receives each distributed weight bucket once and applies it to both the built-in MTP draft runner and the target runner. It supports both SGLang speculative worker layouts, handles SGLang's internal NEXTN-to-EAGLE normalization, and leaves external EAGLE draft models untouched. Draft-weight CPU backup can be enabled so the server remains available while updated weights arrive online.

End-to-end validation was performed with Qwen3.5-2B on Geometry3K GRPO. The training-side MTP weights changed during optimization, and the same updated tensors were loaded into all SGLang draft runners.

Key changes:

  • Add enable_mtp_training and mtp_loss_scaling_factor to the Megatron engine configuration. MTP training implies retaining the model's MTP layers and is incompatible with lm_head_loss_chunk_size.
  • Feed independent MTP labels and loss masks through Megatron forward passes while preserving the main logits-based loss path.
  • Patch Megatron-Core GPTModel and Megatron-Bridge Qwen3VLGPTModel forwarding so Qwen3.5 text and multimodal batches can train MTP.
  • Align MTP labels and masks with padded or packed execution layouts and prevent targets from crossing sequence or padding boundaries.
  • Support MTP training with CP > 1 for wrapper-owned packed THD by applying the same per-sequence zigzag split and rank-local repacking to input IDs, MTP labels, and MTP loss masks.
  • Reuse Megatron-Core's packed, CP-aware rolling semantics to align future-token supervision across CP ranks.
  • Isolate MTP gradients from shared output weights, embeddings, and backbone hidden states.
  • Report the auxiliary mtp_loss in training statistics.
  • Add SGLang speculative-decoding configuration passthrough for NEXTN, speculative steps, EAGLE top-k, draft-token count, external draft-model path, and draft-weight CPU backup.
  • Add an SGLang distributed weight-update bridge that updates both target and built-in MTP draft runners from the same received tensors.
  • Support both SGLang Spec v1 and Spec v2 draft-runner layouts, while failing fast for unsupported SGLang versions, missing draft runners, unsupported load formats, and inference pipeline parallelism.
  • Record rollout/spec_accept_rate and rollout/spec_accept_length from SGLang response metadata.
  • Add a Qwen3.5-2B Geometry3K GRPO example with MTP training and NEXTN rollout enabled.
  • Add unit coverage for padded and packed MTP label/mask layouts, CP zigzag alignment, multimodal forwarding, NEXTN/EAGLE routing, Spec v1/v2 compatibility, and draft/target online weight updates.
  • Regenerate the English and Chinese CLI reference documentation.

Current limitations:

  • MTP training with CP > 1 is supported only for wrapper-owned packed THD. Padded BSHD, VLM, and model-owned THD execution still require CP=1.
  • The current gradient-isolation implementation supports a single MTP prediction layer. Multi-layer MTP gradient propagation is not supported.
  • The SGLang distributed MTP update bridge requires inference pipeline parallel size 1.
  • The compatibility bridge is intentionally pinned to sglang==0.5.10.post1 because it relies on version-specific internal weight-update and draft-runner APIs.

Description

Related Issue

Fixes #(issue)

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 (if applicable; built with ./docs/build_all.sh)
  • Branch is up to date with main
  • Self-reviewed via /review-pr command
  • This PR was created by a coding agent via /create-pr
  • This PR is a breaking change

Breaking Change Details (if applicable):

Additional Context


Need help? Check the
Contributing Guide
or ask in GitHub Discussions!

Train the built-in MTP head jointly with GRPO/SFT on Megatron and use it during rollout through SGLang NEXTN speculative decoding. Online weight synchronization updates both the target model and the built-in MTP draft runner, ensuring that speculative decoding uses the latest RL-trained MTP weights instead of a stale initialization.

MTP is trained through Megatron-Core's auxiliary cross-entropy path with a configurable loss scaling factor of `0.1` by default. AReaL supplies independent MTP label and loss-mask channels while keeping the main forward path logits-based. Shared output weights, backbone hidden states, and embedding inputs are detached from the MTP loss graph, so the backbone receives only the policy/SFT gradient while the MTP-specific parameters learn from future-token supervision.

For packed THD training with context parallelism, AReaL applies the same per-sequence zigzag CP split and rank-local repacking to MTP labels and loss masks as it does to input IDs. Megatron-Core's CP-aware rolling then aligns future-token targets across CP ranks without crossing packed sequence boundaries.

For online rollout, a focused compatibility bridge for `sglang==0.5.10.post1` receives each distributed weight bucket once and applies it to both the built-in MTP draft runner and the target runner. It supports both SGLang speculative worker layouts, handles SGLang's internal `NEXTN`-to-`EAGLE` normalization, and leaves external EAGLE draft models untouched. Draft-weight CPU backup can be enabled so the server remains available while updated weights arrive online.

End-to-end validation was performed with Qwen3.5-2B on Geometry3K GRPO. The training-side MTP weights changed during optimization, and the same updated tensors were loaded into all SGLang draft runners.

Key changes:

- Add `enable_mtp_training` and `mtp_loss_scaling_factor` to the Megatron engine configuration. MTP training implies retaining the model's MTP layers and is incompatible with `lm_head_loss_chunk_size`.
- Feed independent MTP labels and loss masks through Megatron forward passes while preserving the main logits-based loss path.
- Patch Megatron-Core `GPTModel` and Megatron-Bridge `Qwen3VLGPTModel` forwarding so Qwen3.5 text and multimodal batches can train MTP.
- Align MTP labels and masks with padded or packed execution layouts and prevent targets from crossing sequence or padding boundaries.
- Support MTP training with `CP > 1` for wrapper-owned packed THD by applying the same per-sequence zigzag split and rank-local repacking to input IDs, MTP labels, and MTP loss masks.
- Reuse Megatron-Core's packed, CP-aware rolling semantics to align future-token supervision across CP ranks.
- Isolate MTP gradients from shared output weights, embeddings, and backbone hidden states.
- Report the auxiliary `mtp_loss` in training statistics.
- Add SGLang speculative-decoding configuration passthrough for `NEXTN`, speculative steps, EAGLE top-k, draft-token count, external draft-model path, and draft-weight CPU backup.
- Add an SGLang distributed weight-update bridge that updates both target and built-in MTP draft runners from the same received tensors.
- Support both SGLang Spec v1 and Spec v2 draft-runner layouts, while failing fast for unsupported SGLang versions, missing draft runners, unsupported load formats, and inference pipeline parallelism.
- Record `rollout/spec_accept_rate` and `rollout/spec_accept_length` from SGLang response metadata.
- Add a Qwen3.5-2B Geometry3K GRPO example with MTP training and NEXTN rollout enabled.
- Add unit coverage for padded and packed MTP label/mask layouts, CP zigzag alignment, multimodal forwarding, NEXTN/EAGLE routing, Spec v1/v2 compatibility, and draft/target online weight updates.
- Regenerate the English and Chinese CLI reference documentation.

Current limitations:

- MTP training with `CP > 1` is supported only for wrapper-owned packed THD. Padded BSHD, VLM, and model-owned THD execution still require `CP=1`.
- The current gradient-isolation implementation supports a single MTP prediction layer. Multi-layer MTP gradient propagation is not supported.
- The SGLang distributed MTP update bridge requires inference pipeline parallel size `1`.
- The compatibility bridge is intentionally pinned to `sglang==0.5.10.post1` because it relies on version-specific internal weight-update and draft-runner APIs.
@HT-Yuan

HT-Yuan commented Sep 1, 2026 •

Copy link
Copy Markdown
Collaborator Author

To make testing and review easier, the work from #1650 has been moved to #1659 .

I think there are still two valuable follow-up directions:

  1. Broader context-parallel support. feat: add MTP training with speculative decoding rollout #1659 currently supports MTP with CP only for the wrapper-owned packed THD path. Fully supporting Qwen3.5 padded BSHD, VLM, and model-owned layouts requires more than splitting mtp_labels and mtp_loss_mask: the input and embedding layout, position information, attention/GDN CP semantics, multimodal embeddings, and output/loss reconstruction all need to remain consistent across CP ranks. I intentionally left this out because it is a separate model-level integration and would make the current PR much broader and harder to validate.

  2. Correct gradient isolation for multi-depth MTP training. The current detach behavior is correct for a single MTP prediction layer, which is the configuration used by the Qwen3.5 model we tested. With multiple MTP depths, however, detaching the hidden states at every depth breaks the gradient path from later MTP losses back into earlier MTP blocks. A follow-up could detach only at the backbone-to-MTP boundary, while preserving the computation graph between MTP depths and continuing to isolate the shared embedding and output weights. This was not included because multi-depth MTP is not required by the current target model and needs dedicated loss, gradient, and optimizer-update equivalence tests.

@HT-Yuan

HT-Yuan commented Sep 1, 2026 •

Copy link
Copy Markdown
Collaborator Author
Clipboard_Screenshot_1788273174 Clipboard_Screenshot_1788273311

@sitabulaixizawaluduo sitabulaixizawaluduo added safe-to-test Ready to run unit-tests in a PR. and removed safe-to-test Ready to run unit-tests in a PR. labels Sep 1, 2026
Comment thread areal/engine/megatron_utils/megatron_bridge_patches.py Outdated
Comment thread areal/engine/megatron_utils/megatron_bridge_patches.py Outdated
huaqingyuan and others added 3 commits September 3, 2026 03:45
Keep next-token-aligned masks unchanged before MCore's per-layer
roll, and detach internal output-layer weights for untied models.

Co-authored-by: Cursor <cursoragent@cursor.com>
Resolve generated CLI reference conflicts by regenerating documentation
from the combined configuration definitions.

Co-authored-by: Cursor <cursoragent@cursor.com>
Let MCore apply the optimizer loss scale to both the main backward
path and separately seeded MTP/MoE auxiliary gradients. This prevents
FP16 optimizer unscaling from suppressing auxiliary updates.

Co-authored-by: Cursor <cursoragent@cursor.com>
Comment thread areal/engine/megatron_engine.py
@kkli08

kkli08 commented Sep 3, 2026

Copy link
Copy Markdown

To make testing and review easier, the work from #1650 has been moved to #1659 .

I think there are still two valuable follow-up directions:

  1. Broader context-parallel support. feat: add MTP training with speculative decoding rollout #1659 currently supports MTP with CP only for the wrapper-owned packed THD path. Fully supporting Qwen3.5 padded BSHD, VLM, and model-owned layouts requires more than splitting mtp_labels and mtp_loss_mask: the input and embedding layout, position information, attention/GDN CP semantics, multimodal embeddings, and output/loss reconstruction all need to remain consistent across CP ranks. I intentionally left this out because it is a separate model-level integration and would make the current PR much broader and harder to validate.
  2. Correct gradient isolation for multi-depth MTP training. The current detach behavior is correct for a single MTP prediction layer, which is the configuration used by the Qwen3.5 model we tested. With multiple MTP depths, however, detaching the hidden states at every depth breaks the gradient path from later MTP losses back into earlier MTP blocks. A follow-up could detach only at the backbone-to-MTP boundary, while preserving the computation graph between MTP depths and continuing to isolate the shared embedding and output weights. This was not included because multi-depth MTP is not required by the current target model and needs dedicated loss, gradient, and optimizer-update equivalence tests.

LGTM! I see these as two separate follow-ups:

  1. Multi-depth MTP gradient isolation: I’ve prepared a small PR on top of feat: add MTP training with speculative decoding rollout #1659, with exact gradient and optimizer-update tests for distinct and repeated depths.
  2. Broader Qwen3.5 GDN + MTP + CP support: This will remain a separate follow-up. AReaL currently pins Megatron-Core 0.17.0; GDN CP and packed-sequence support first shipped in MCore 0.18.0, while the remaining packed-THD MTP+CP correctness fix is still pending in fix(mtp): use padded cu_seqlens in MTP roll for THD with CP NVIDIA/Megatron-LM#4495.
    cc @sitabulaixizawaluduo

@sitabulaixizawaluduo

Copy link
Copy Markdown
Collaborator

To make testing and review easier, the work from #1650 has been moved to #1659 .
I think there are still two valuable follow-up directions:

  1. Broader context-parallel support. feat: add MTP training with speculative decoding rollout #1659 currently supports MTP with CP only for the wrapper-owned packed THD path. Fully supporting Qwen3.5 padded BSHD, VLM, and model-owned layouts requires more than splitting mtp_labels and mtp_loss_mask: the input and embedding layout, position information, attention/GDN CP semantics, multimodal embeddings, and output/loss reconstruction all need to remain consistent across CP ranks. I intentionally left this out because it is a separate model-level integration and would make the current PR much broader and harder to validate.
  2. Correct gradient isolation for multi-depth MTP training. The current detach behavior is correct for a single MTP prediction layer, which is the configuration used by the Qwen3.5 model we tested. With multiple MTP depths, however, detaching the hidden states at every depth breaks the gradient path from later MTP losses back into earlier MTP blocks. A follow-up could detach only at the backbone-to-MTP boundary, while preserving the computation graph between MTP depths and continuing to isolate the shared embedding and output weights. This was not included because multi-depth MTP is not required by the current target model and needs dedicated loss, gradient, and optimizer-update equivalence tests.

LGTM! I see these as two separate follow-ups:

  1. Multi-depth MTP gradient isolation: I’ve prepared a small PR on top of feat: add MTP training with speculative decoding rollout #1659, with exact gradient and optimizer-update tests for distinct and repeated depths.
  2. Broader Qwen3.5 GDN + MTP + CP support: This will remain a separate follow-up. AReaL currently pins Megatron-Core 0.17.0; GDN CP and packed-sequence support first shipped in MCore 0.18.0, while the remaining packed-THD MTP+CP correctness fix is still pending in fix(mtp): use padded cu_seqlens in MTP roll for THD with CP NVIDIA/Megatron-LM#4495.
    cc @sitabulaixizawaluduo

It's a valuable piece of work.BTW. The CP for Qwen3.5 is already available in the branch feature/qwen35-vlm-thd-cp. You need to manually upgrade to MCore==0.18.2 and megatron-bridge==0.5.1. Due to version conflicts, this part cannot be merged into the main branch at the moment, but the subsequent support for MTP CP can be referred to from this implementation.

Comment thread areal/models/mcore/registry.py
Fail before model construction when MTP training requests more than
one prediction layer, whose gradients are not fully supported yet.

Co-authored-by: Cursor <cursoragent@cursor.com>

@sitabulaixizawaluduo sitabulaixizawaluduo left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

LGTM

@sitabulaixizawaluduo
sitabulaixizawaluduo merged commit 7823be6 into main Sep 6, 2026
6 checks passed
@sitabulaixizawaluduo
sitabulaixizawaluduo deleted the feat/megatron-mtp-cp branch September 6, 2026 03:39
Vivicai1005 pushed a commit to Vivicai1005/AReaL that referenced this pull request Sep 8, 2026
…ct#1659)

* feat: add MTP training with NEXTN speculative decoding rollout

Train the built-in MTP head jointly with GRPO/SFT on Megatron and use it during rollout through SGLang NEXTN speculative decoding. Online weight synchronization updates both the target model and the built-in MTP draft runner, ensuring that speculative decoding uses the latest RL-trained MTP weights instead of a stale initialization.

MTP is trained through Megatron-Core's auxiliary cross-entropy path with a configurable loss scaling factor of `0.1` by default. AReaL supplies independent MTP label and loss-mask channels while keeping the main forward path logits-based. Shared output weights, backbone hidden states, and embedding inputs are detached from the MTP loss graph, so the backbone receives only the policy/SFT gradient while the MTP-specific parameters learn from future-token supervision.

For packed THD training with context parallelism, AReaL applies the same per-sequence zigzag CP split and rank-local repacking to MTP labels and loss masks as it does to input IDs. Megatron-Core's CP-aware rolling then aligns future-token targets across CP ranks without crossing packed sequence boundaries.

For online rollout, a focused compatibility bridge for `sglang==0.5.10.post1` receives each distributed weight bucket once and applies it to both the built-in MTP draft runner and the target runner. It supports both SGLang speculative worker layouts, handles SGLang's internal `NEXTN`-to-`EAGLE` normalization, and leaves external EAGLE draft models untouched. Draft-weight CPU backup can be enabled so the server remains available while updated weights arrive online.

End-to-end validation was performed with Qwen3.5-2B on Geometry3K GRPO. The training-side MTP weights changed during optimization, and the same updated tensors were loaded into all SGLang draft runners.

Key changes:

- Add `enable_mtp_training` and `mtp_loss_scaling_factor` to the Megatron engine configuration. MTP training implies retaining the model's MTP layers and is incompatible with `lm_head_loss_chunk_size`.
- Feed independent MTP labels and loss masks through Megatron forward passes while preserving the main logits-based loss path.
- Patch Megatron-Core `GPTModel` and Megatron-Bridge `Qwen3VLGPTModel` forwarding so Qwen3.5 text and multimodal batches can train MTP.
- Align MTP labels and masks with padded or packed execution layouts and prevent targets from crossing sequence or padding boundaries.
- Support MTP training with `CP > 1` for wrapper-owned packed THD by applying the same per-sequence zigzag split and rank-local repacking to input IDs, MTP labels, and MTP loss masks.
- Reuse Megatron-Core's packed, CP-aware rolling semantics to align future-token supervision across CP ranks.
- Isolate MTP gradients from shared output weights, embeddings, and backbone hidden states.
- Report the auxiliary `mtp_loss` in training statistics.
- Add SGLang speculative-decoding configuration passthrough for `NEXTN`, speculative steps, EAGLE top-k, draft-token count, external draft-model path, and draft-weight CPU backup.
- Add an SGLang distributed weight-update bridge that updates both target and built-in MTP draft runners from the same received tensors.
- Support both SGLang Spec v1 and Spec v2 draft-runner layouts, while failing fast for unsupported SGLang versions, missing draft runners, unsupported load formats, and inference pipeline parallelism.
- Record `rollout/spec_accept_rate` and `rollout/spec_accept_length` from SGLang response metadata.
- Add a Qwen3.5-2B Geometry3K GRPO example with MTP training and NEXTN rollout enabled.
- Add unit coverage for padded and packed MTP label/mask layouts, CP zigzag alignment, multimodal forwarding, NEXTN/EAGLE routing, Spec v1/v2 compatibility, and draft/target online weight updates.
- Regenerate the English and Chinese CLI reference documentation.

Current limitations:

- MTP training with `CP > 1` is supported only for wrapper-owned packed THD. Padded BSHD, VLM, and model-owned THD execution still require `CP=1`.
- The current gradient-isolation implementation supports a single MTP prediction layer. Multi-layer MTP gradient propagation is not supported.
- The SGLang distributed MTP update bridge requires inference pipeline parallel size `1`.
- The compatibility bridge is intentionally pinned to `sglang==0.5.10.post1` because it relies on version-specific internal weight-update and draft-runner APIs.

* test: use HttpGenerationResult in rollout version race test

* fix(engine): align MTP masks and detach untied output weights

Keep next-token-aligned masks unchanged before MCore's per-layer
roll, and detach internal output-layer weights for untied models.

Co-authored-by: Cursor <cursoragent@cursor.com>

* fix(engine): unify Megatron main and auxiliary loss scaling

Let MCore apply the optimizer loss scale to both the main backward
path and separately seeded MTP/MoE auxiliary gradients. This prevents
FP16 optimizer unscaling from suppressing auxiliary updates.

Co-authored-by: Cursor <cursoragent@cursor.com>

* fix(models): reject unsupported multilayer MTP training

Fail before model construction when MTP training requests more than
one prediction layer, whose gradients are not fully supported yet.

Co-authored-by: Cursor <cursoragent@cursor.com>

---------

Co-authored-by: huaqingyuan <huaqingyuan@tencent.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

safe-to-test Ready to run unit-tests in a PR.

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants