Skip to content

feat(megatron): add MTP auxiliary training - #1653

Draft
kkli08 wants to merge 1 commit into
areal-project:mainfrom
kkli08:kkli08/mtp-training-support
Draft

kkli08 wants to merge 1 commit into
areal-project:mainfrom
kkli08:kkli08/mtp-training-support

Conversation

@kkli08

@kkli08 kkli08 commented Aug 31, 2026

Copy link
Copy Markdown

Description

Adds opt-in auxiliary Multi-Token Prediction (MTP) training for
Megatron-Bridge actors while preserving AReaL's existing logits-return
contract.

The feature is disabled by default through enable_mtp_training=False and
uses mtp_loss_scaling_factor=0.1 unless configured otherwise.

What changes

  • Builds boundary-safe next-token targets and reuses the actor loss_mask, so
    MTP targets do not cross packed trajectory boundaries.
  • Injects MTP supervision without passing labels through GPTModel, keeping
    AReaL's policy-loss path based on returned logits unchanged.
  • Isolates auxiliary gradients from the backbone, shared embedding, and LM
    head while preserving gradients across distinct and repeated MTP depths.
  • Applies the same microbatch/global weighting as the main objective, including
    CP-global token normalization for wrapper-packed THD inputs.
  • Skips MTP computation for eval and other logits-only forwards.
  • Adds explicit validation, per-depth MTP metrics, CLI documentation, and
    focused regression coverage.

Support matrix

Path Status
Text-only Megatron-Bridge actor, CP=1 Implemented and unit-tested
Qwen3.5/GDN padded BSHD path CP=1 only
Wrapper-owned packed THD with CP-aware MCore CP path implemented and unit-tested; real multi-rank GPU validation pending
Distinct or repeated multi-depth MTP Gradient chaining covered by exact-gradient tests

Related Issue

Related to #1445. This is an independent implementation focused on AReaL's
existing logits-based actor path, gradient isolation, and packed context
parallelism.

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): N/A

Additional Context

Validation completed

  • SKIP=generate-cli-docs pre-commit run --all-files (the same command used by
    the current pre-commit CI workflow)
  • pytest -q tests/test_mcore_mtp_training.py — 20 passed
  • Conventional Commit hook
  • git diff --check

The focused CPU tests mock distributed collectives and model hooks. They cover
trajectory boundaries, CP layout and loss scaling, distinct/repeated-depth
gradients, backbone/embedding/LM-head isolation, eval bypass, runtime capability
checks, and checkpoint argument handling.

Not yet run

  • End-to-end Linux GPU training with the public Megatron/Megatron-Bridge stack
  • Real multi-rank CP=2 training and CP=1/CP=2 gradient parity
  • Full ./docs/build_all.sh documentation build
  • The config test in tests/test_megatron_lm_head.py locally, because the
    pinned Megatron-Core dependency is Linux-only; it is included for Linux CI

Current limitations

  • Actor training only; critic objectives are rejected.
  • Text-only; multimodal batches are rejected.
  • Tree training and model-owned THD are not supported.
  • The pinned Qwen3.5/GDN padded path remains CP=1.
  • calculate_per_token_loss=True is not supported with MTP context
    parallelism.

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

Add opt-in auxiliary MTP training for Megatron-Bridge actors while
preserving AReaL's logits-return contract.

Key changes:
- build boundary-safe targets and apply actor loss masks
- isolate auxiliary gradients while preserving multi-depth chains
- normalize wrapper-packed THD loss across context-parallel ranks
- add config validation, CLI docs, and focused regression tests

Refs: areal-project#1445
Signed-off-by: Ke Li <77185597+kkli08@users.noreply.github.com>
@sitabulaixizawaluduo

Copy link
Copy Markdown
Collaborator

Thank you for your contribution, PR1652 is already working on similar issues. Could you describe the differences between your work and his? If they are of the same type, you can comment on the former.

@kkli08

kkli08 commented Sep 1, 2026

Copy link
Copy Markdown
Author

Thank you for your contribution, PR1652 is already working on similar issues. Could you describe the differences between your work and his? If they are of the same type, you can comment on the former.

Thanks for pointing this out — I assume you meant #1650. After comparing the two, there is substantial overlap, and #1650 covers a broader end-to-end scope.
The complementary pieces from #1653 are:

  • MTP training with wrapper-packed THD context parallelism; feat: add MTP training with speculative decoding rollout #1650 currently limits it to CP=1
  • MTP loss scaling aligned with AReaL’s global microbatch weighting and CP normalization
  • Multi-depth gradient-chain correctness, together with focused eval, checkpoint, runtime-compatibility, and regression tests

Part of my previous work has also related with MTP and AReaL, including long-context performance evaluation with MTP On/Off, auxiliary-training tuning, and rollout-side integration. I'd love to turn the generally useful lessons from that experience into public and reproducible open-source contributions.
I’m happy to use #1650 as the base and contribute these complementary pieces there.

@HT-Yuan would you be open to collaborating on the integration and possible follow-up work? : )

@kkli08

kkli08 commented Sep 1, 2026

Copy link
Copy Markdown
Author

Thank you for your contribution, PR1652 is already working on similar issues. Could you describe the differences between your work and his? If they are of the same type, you can comment on the former.

Thanks for pointing this out — I assume you meant #1650. After comparing the two, there is substantial overlap, and #1650 covers a broader end-to-end scope. The complementary pieces from #1653 are:

  • MTP training with wrapper-packed THD context parallelism; feat: add MTP training with speculative decoding rollout #1650 currently limits it to CP=1
  • MTP loss scaling aligned with AReaL’s global microbatch weighting and CP normalization
  • Multi-depth gradient-chain correctness, together with focused eval, checkpoint, runtime-compatibility, and regression tests

Part of my previous work has also related with MTP and AReaL, including long-context performance evaluation with MTP On/Off, auxiliary-training tuning, and rollout-side integration. I'd love to turn the generally useful lessons from that experience into public and reproducible open-source contributions. I’m happy to use #1650 as the base and contribute these complementary pieces there.

@HT-Yuan would you be open to collaborating on the integration and possible follow-up work? : )

From a previous end-to-end math RL run with a 32K context limit. Here, K2D1 means rollout K=2 with one trainable MTP depth (D=1).
rollout

@sitabulaixizawaluduo

Copy link
Copy Markdown
Collaborator

Thank you for your contribution, PR1652 is already working on similar issues. Could you describe the differences between your work and his? If they are of the same type, you can comment on the former.

Thanks for pointing this out — I assume you meant #1650. After comparing the two, there is substantial overlap, and #1650 covers a broader end-to-end scope. The complementary pieces from #1653 are:

  • MTP training with wrapper-packed THD context parallelism; feat: add MTP training with speculative decoding rollout #1650 currently limits it to CP=1
  • MTP loss scaling aligned with AReaL’s global microbatch weighting and CP normalization
  • Multi-depth gradient-chain correctness, together with focused eval, checkpoint, runtime-compatibility, and regression tests

Part of my previous work has also related with MTP and AReaL, including long-context performance evaluation with MTP On/Off, auxiliary-training tuning, and rollout-side integration. I'd love to turn the generally useful lessons from that experience into public and reproducible open-source contributions. I’m happy to use #1650 as the base and contribute these complementary pieces there.
@HT-Yuan would you be open to collaborating on the integration and possible follow-up work? : )

From a previous end-to-end math RL run with a 32K context limit. Here, K2D1 means rollout K=2 with one trainable MTP depth (D=1). rollout

It seems that the acceleration here is starting to converge later than at the end, what is the reason for this?

@kkli08

kkli08 commented Sep 1, 2026 •

Copy link
Copy Markdown
Author

Thank you for your contribution, PR1652 is already working on similar issues. Could you describe the differences between your work and his? If they are of the same type, you can comment on the former.

Thanks for pointing this out — I assume you meant #1650. After comparing the two, there is substantial overlap, and #1650 covers a broader end-to-end scope. The complementary pieces from #1653 are:

  • MTP training with wrapper-packed THD context parallelism; feat: add MTP training with speculative decoding rollout #1650 currently limits it to CP=1
  • MTP loss scaling aligned with AReaL’s global microbatch weighting and CP normalization
  • Multi-depth gradient-chain correctness, together with focused eval, checkpoint, runtime-compatibility, and regression tests

Part of my previous work has also related with MTP and AReaL, including long-context performance evaluation with MTP On/Off, auxiliary-training tuning, and rollout-side integration. I'd love to turn the generally useful lessons from that experience into public and reproducible open-source contributions. I’m happy to use #1650 as the base and contribute these complementary pieces there.
@HT-Yuan would you be open to collaborating on the integration and possible follow-up work? : )

From a previous end-to-end math RL run with a 32K context limit. Here, K2D1 means rollout K=2 with one trainable MTP depth (D=1). rollout

It seems that the acceleration here is starting to converge later than at the end, what is the reason for this?

Good observation. This is raw rollout time from three separate online RL runs using Qwen3.5-35B-A3B, not token-normalized throughput. As the policies evolve, their output lengths and EOS behavior diverge;

The later convergence likely reflects less decode work and a larger share of fixed overhead. I think MTP is not losing effectiveness, the K2D1 accepted length still reached around 2.46.

spec_accept_length

@HT-Yuan

HT-Yuan commented Sep 1, 2026

Copy link
Copy Markdown
Collaborator

Thank you for your contribution, PR1652 is already working on similar issues. Could you describe the differences between your work and his? If they are of the same type, you can comment on the former.

Thanks for pointing this out — I assume you meant #1650. After comparing the two, there is substantial overlap, and #1650 covers a broader end-to-end scope. The complementary pieces from #1653 are:

  • MTP training with wrapper-packed THD context parallelism; feat: add MTP training with speculative decoding rollout #1650 currently limits it to CP=1
  • MTP loss scaling aligned with AReaL’s global microbatch weighting and CP normalization
  • Multi-depth gradient-chain correctness, together with focused eval, checkpoint, runtime-compatibility, and regression tests

Part of my previous work has also related with MTP and AReaL, including long-context performance evaluation with MTP On/Off, auxiliary-training tuning, and rollout-side integration. I'd love to turn the generally useful lessons from that experience into public and reproducible open-source contributions. I’m happy to use #1650 as the base and contribute these complementary pieces there.

@HT-Yuan would you be open to collaborating on the integration and possible follow-up work? : )

@kkli08
Absolutely, I’d be very happy to collaborate! 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 CP 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.

Feel free to share your thoughts. cc @sitabulaixizawaluduo

@github-actions

Copy link
Copy Markdown

This pull request has been automatically marked as stale because it has not had recent activity within the last 14 days.

Please add a comment or push new commits to keep it active.

Thank you for your contribution!

@github-actions github-actions Bot added the stale label Sep 16, 2026

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

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants