Add Megatron-Core MTP training for Qwen3.5 MoE - #198
Closed
liuzijing2014 wants to merge 1 commit into
Closed
liuzijing2014 wants to merge 1 commit into
liuzijing2014 wants to merge 1 commit into
Conversation
Train a model's own MTP layer with Megatron-Core's native MultiTokenPredictionLayer, MoE and expert-parallel token dispatch, while TorchSpec keeps the data, loss, optimizer and checkpoints. - torchspec.megatron: TP1/PP1 process groups with any EP size, and gradient sync/clipping that follows Megatron DDP ownership for replicated vs expert-sharded parameters. - torchspec.models.qwen3_5_mtp: build the Qwen3.5 MoE MTP layer from native Megatron modules and import it from a Hugging Face checkpoint (zero-centered RMSNorm, gated grouped QKV, routed and shared experts), returning a per-tensor import report. - torchspec.models.target_hidden: capture both sides of the target's final norm. The MTP consumes the post-norm state, as Megatron GPT+MTP and vLLM do. - torchspec.training.mtp: hard-label CE against x[t+2], with a global token normalizer for gradient accumulation. Signed-off-by: Zijing Liu <liuzijing2014@gmail.com>
liuzijing2014
force-pushed
the
feat/qwen35-megatron-mtp
branch
from
September 25, 2026 00:03
bb2cb59 to
10c97bc
Compare
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Purpose
Qwen3.5 ships with a trained MTP layer, but TorchSpec has no way to fine-tune a model's own MTP head today. This PR adds a small Megatron-Core integration so that layer can be trained with Megatron's native
MultiTokenPredictionLayer, MoE and expert-parallel token dispatch. TorchSpec keeps owning the data, loss, optimizer and checkpoints. The first target is Qwen3.5-35B-A3B at TP1/PP1 with EP1 or EP4.What's in it
torchspec/megatron.py: sets up TP1/PP1 process groups for any EP size and seeds Megatron's CUDA RNG right after group creation. Every rank in an EP group gets the same batch, so a routed expert's gradient already sums one copy of each token per EP rank.synchronize_gradientsandclip_grad_norm_therefore follow Megatron DDP: both replicated and expert gradients are scaled by 1/dense-DP, and each is reduced over its own group.torchspec/models/qwen3_5_mtp.py: builds the MTP layer from native Megatron modules (local spec) and imports it from the HF checkpoint. The import:1 + w;[query, gate]rows into Megatron's per-group[queries, gates, key, value]QKV layout;Parameters are NaN-filled before loading. The loader returns a per-tensor mapping report and fails on anything left unloaded or unconsumed. The router runs without an aux loss, so the objective is plain cross entropy.
torchspec/models/target_hidden.py: hooks both sides of the target's final RMSNorm. The MTP consumes the post-norm state, which is also what Megatron's GPT+MTP path and vLLM's Qwen3.5 MTP use.torchspec/training/mtp.py: hard-labelF.cross_entropy. MTP rowt(inputsh_tandx[t+1]) is scored againstx[t+2]; masked targets and the wrapped tail are excluded. An optional normalizer keeps the step loss a global token mean under gradient accumulation.New
megatronextra:megatron-core==0.18.0(needs Python 3.12).Test plan
Setup:
Qwen/Qwen3.5-35B-A3B@59d61f3. Shard sha256 values were checked against the Hub LFS oids before and after the runs.Aeala/ShareGPT_Vicuna_unfiltered@8b0048a, rendered with TorchSpec'sqwentemplate, max 2048 tokens. That is 2.31M tokens and 1.71M supervised targets.pytest tests/test_mtp.py tests/test_target_hidden.pyThe end-to-end checks below used a one-epoch training harness built on these APIs; the harness is not part of this PR.
Qwen3_5MoeDecoderLayeron 64 samples.Results
Unit tests: 5 passed.
F.cross_entropyrecomputeNotes
SequentialMLP). The MTP layer runs in fp32 and isn't tuned for speed: about 7.5 s/step at EP1 and 3.0 s/step at EP4 for 16 sequences of up to 2k tokens, target forward included.