Skip to content

Add Megatron-Core MTP training for Qwen3.5 MoE - #198

Closed
liuzijing2014 wants to merge 1 commit into
lightseekorg:mainfrom
liuzijing2014:feat/qwen35-megatron-mtp
Closed

liuzijing2014 wants to merge 1 commit into
lightseekorg:mainfrom
liuzijing2014:feat/qwen35-megatron-mtp

Conversation

@liuzijing2014

@liuzijing2014 liuzijing2014 commented Sep 24, 2026 •

Copy link
Copy Markdown
Contributor

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_gradients and clip_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:

    • converts Qwen's zero-centered RMSNorm weights to 1 + w;
    • reorders the per-head [query, gate] rows into Megatron's per-group [queries, gates, key, value] QKV layout;
    • loads the routed experts, shared expert and shared-expert gate.

    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-label F.cross_entropy. MTP row t (inputs h_t and x[t+1]) is scored against x[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 megatron extra: megatron-core==0.18.0 (needs Python 3.12).

Test plan

Setup:

  • Model: Qwen/Qwen3.5-35B-A3B @ 59d61f3. Shard sha256 values were checked against the Hub LFS oids before and after the runs.
  • Data: 2048 conversations from Aeala/ShareGPT_Vicuna_unfiltered @ 8b0048a, rendered with TorchSpec's qwen template, max 2048 tokens. That is 2.31M tokens and 1.71M supervised targets.
  • Training: global batch 16, so 128 steps per epoch. AdamW at 2e-5 with warmup, then cosine decay. MTP weights in fp32; target frozen in bf16.
  • Environment: one GB200 node; torch 2.11 / CUDA 13.0, megatron-core 0.18.0, transformers 5.12.1; deterministic algorithms on, TF32 off.
  1. pytest tests/test_mtp.py tests/test_target_hidden.py

The end-to-end checks below used a one-epoch training harness built on these APIs; the harness is not part of this PR.

  1. HF parity: compares the Megatron MTP against a reference assembled from HF's own Qwen3_5MoeDecoderLayer on 64 samples.
  2. One epoch at EP1 and at EP4:
    • Preflight: one step that checks the native modules, the import gate, that the MTP input is exactly the final-norm output, that grads are finite and non-zero, that the target is untouched, that the update was applied, and that ranks agree. CE is then recomputed from the saved logits/labels.
    • Reference segment: train from the HF weights through step 64, checkpoint, then run step 65.
    • Resume: restart from the step-64 checkpoint, require step 65 to match the reference segment, then finish the epoch.
    • Reload: load the end-of-epoch checkpoint in a fresh process.

Results

Unit tests: 5 passed.

Check EP1 EP4
Native Megatron MTP/MoE modules pass pass, 64 experts per rank
HF import 528 tensors, all exact; 0 missing / 0 unexpected keys (vision tower allowlisted) 4 × 144 tensors, all exact; each of the 256 experts owned by exactly one rank
MTP input is the final-norm output, checked every micro-batch pass pass
CE vs F.cross_entropy recompute 2.8e-8 relative 4.2e-8 relative
Step 1 loss / grad norm 0.91130 / 3.4275 0.91130 / 3.4275
Resume at step 64, compare step 65 bit-for-bit identical bit-for-bit identical on all 4 ranks
Epoch accounting 128 steps, all 2048 samples exactly once, 2,314,964 tokens same
End-of-epoch checkpoint reload exact (params, Adam state, RNG, probe loss) exact
  • HF parity (64 samples, 52,232 targets): argmax agreement is 100%, and mean CE is 0.825344 for both. Over the 66,310 of 66,311 rows where both routers pick the same experts, the max |Δ| on the final hidden state is 6.2e-5. The one remaining row is a top-8 near-tie (8th and 9th router logits 1.9e-6 apart) where the two pick a different expert.
  • Training: loss goes from 0.911 to 0.628 over the epoch. EP1 and EP4 see the same samples at every step, and their per-step losses stay within 1e-4 of each other.
  • Why post-norm input: the imported MTP scores CE 0.825 (77.4% top-1) on the post-norm target state versus 0.890 (76.3%) on the pre-norm state, which fits it having been trained on post-norm states.

Notes

  • TP1/PP1 only, and no Transformer Engine yet (Megatron local spec with 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.
  • Every rank in an EP group has to get the same batch; shard data by expert-data-parallel rank.
  • This isn't wired into the Ray/vLLM pipeline yet. In the example, the target runs through HF transformers.

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
liuzijing2014 force-pushed the feat/qwen35-megatron-mtp branch from bb2cb59 to 10c97bc Compare September 25, 2026 00:03
@liuzijing2014
liuzijing2014 deleted the feat/qwen35-megatron-mtp branch September 25, 2026 06:19
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.

1 participant