Repository navigation
Conversation
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>
|
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.
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. @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). |
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.
|
@kkli08 I think there are still two valuable follow-up directions:
Feel free to share your thoughts. cc @sitabulaixizawaluduo |
|
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! |



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=Falseanduses
mtp_loss_scaling_factor=0.1unless configured otherwise.What changes
loss_mask, soMTP targets do not cross packed trajectory boundaries.
GPTModel, keepingAReaL's policy-loss path based on returned logits unchanged.
head while preserving gradients across distinct and repeated MTP depths.
CP-global token normalization for wrapper-packed THD inputs.
focused regression coverage.
Support matrix
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
Checklist
Contributing Guide
pre-commit run --all-files)./docs/build_all.sh)main/review-prcommand/create-prBreaking Change Details (if applicable): N/A
Additional Context
Validation completed
SKIP=generate-cli-docs pre-commit run --all-files(the same command used bythe current pre-commit CI workflow)
pytest -q tests/test_mcore_mtp_training.py— 20 passedgit diff --checkThe 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
./docs/build_all.shdocumentation buildtests/test_megatron_lm_head.pylocally, because thepinned Megatron-Core dependency is Linux-only; it is included for Linux CI
Current limitations
calculate_per_token_loss=Trueis not supported with MTP contextparallelism.
Need help? Check the
Contributing Guide
or ask in GitHub Discussions!