Skip to content

(fsdp2 dev.) feat(moe): expert parallelism for MoE drafters on the FSDP2 backend (stacked on #915 + #931) - #932

Draft
yushengsu-thu wants to merge 12 commits into
sgl-project:mainfrom
yushengsu-thu:moe-ep-fsdp2
Draft

yushengsu-thu wants to merge 12 commits into
sgl-project:mainfrom
yushengsu-thu:moe-ep-fsdp2

Conversation

@yushengsu-thu

Copy link
Copy Markdown

Draft, stacked on #915 (FSDP2 backend) and #931 (@sherlockwu's MoE stack #811/#812 rebased on main). This branch is moe-stack-rebased (#931) + a merge of codex/fsdp2-backend (#915; one conflict in training/backend.py, resolved by applying the MoE checkpoint-naming converters around the FSDP2 state-dict seam) + one feature commit. Until those land, review the last commit: git diff 7965b568..12c3f9d2.

Motivation

The MoE stack trains routed experts as ordinary FSDP parameters, so every micro-batch re-gathers all of them: for the DeepSeek-V4-Flash drafter that is 32.7B expert parameters (~65 GB per micro-batch, 46 s/step on 8x B200 per the stack's last commit), which is why that commit freezes warm-started experts. #882/#883 add expert parallelism on the FSDP1 side, but only with NO_SHARD, and carry their own grad-norm partitioning, optimizer owner-split and per-rank checkpoint shards.

This PR adds expert parallelism on the FSDP2 backend so FSDP and EP compose, and the optimizer, grad-norm and checkpoint paths stay exactly as they are.

What this adds

  • training.expert_parallel_size (FSDP2 only; requires a sharded fsdp_sharding, tp_size=1, no sequence parallelism, a world size divisible by it, and n_routed_experts divisible by it).
  • init_distributed builds a 2-D (efsdp, ep) mesh (ep = consecutive ranks, one node for ep <= GPUs per node); ParallelConfig carries it.
  • specforge/modeling/draft/moe/expert_parallel.py: the EP layout, shard_expert_parameter (the stacked [E, ...] weights become DTensor Shard(0) over ep), and the differentiable seams gather_tokens (all-gather, backward reduce-scatter) and scatter_outputs (reduce-scatter, backward all-gather).
  • MoELayer.forward under EP: all-gather the group's tokens, run the replicated router on them, compute only the locally owned experts for all of them, reduce-scatter the fp32 partial outputs, add the shared expert on the local tokens. The EP axis is carved out of data parallelism, so ranks keep distinct data and the dense part of the draft is not computed twice.
  • GroupedExperts.apply_expert_parallel and a local-slice dispatch for both sorted_loop and grouped_mm (sorting by expert makes this rank's experts one contiguous run of slots). Zero-valued terms keep the gathered input, combine weights and expert parameters on the graph when a rank receives no token, so every rank issues the same collectives and FSDP2 sees the same gradient set.
  • FSDP2TrainingBackend: applies EP before wrapping; a per-parameter shard_placement_fn (ShardPlacementResult, available in torch 2.13) puts expert slices on the efsdp mesh and everything else on the full mesh; rescales expert gradients by 1/ep at the optimizer boundary; get_model_state_dict(full_state_dict=True) still gathers the full [E, ...] tensors, so the MoE naming converters, resume and export are unchanged.
  • Guards: the FSDP1 backend and NO_SHARD reject EP with explicit errors.
  • Recipe examples/configs/online/disaggregated/external/deepseek-v4-flash-dspark-moe-ep4-disaggregated.yaml (the MoE arm with backend: fsdp2, expert_parallel_size: 4: the DP4 trainer is one EP group, 16 of 64 experts per rank, experts trained) and docs (customization guide, training guide, recipe runbook, examples README, moe/DESIGN.md).

Why no owner-split, per-rank shards or norm partitioning

With expert slices as DTensors over (efsdp, ep), every parameter element lives on exactly one rank: the optimizer's local-shard sum of squares all-reduced over WORLD is already the exact global norm, Adam state is per shard, and DCP gathers full tensors on rank 0. Gradient scale: FSDP2 averages dense gradients over the world and expert gradients over efsdp = world / ep, but an expert owner's gradient already sums its whole group's tokens, so the backend multiplies expert gradients by 1/ep. The 4-rank test checks every gradient against a plain data-parallel reference.

Tests

CPU gloo, torch 2.13.0 / transformers 5.12.1 (no GPU on this machine):

  • tests/test_modeling/test_moe_expert_parallel.py: 2-rank layer parity (outputs, input gradients, expert-slice gradients, router and shared-expert partial gradients summing to the reference, routing counts over the gathered tokens, official-naming checkpoint); a rank with no selected expert still completes every collective; 4-rank FSDP2TrainingBackend with ep=2 (efsdp=2) and ep=4 (efsdp=1): all gradients vs the data-parallel reference, full checkpoint in official naming on rank 0, zero-then-reload round trip through set_model_state_dict; the unsharded layer is unchanged.
  • tests/test_config/test_schema.py: the new field's validators and world-size check.
python -m unittest tests.test_modeling.test_moe_expert_parallel tests.test_config.test_schema \
  tests.test_modeling.test_moe_deepseek_v4 tests.test_config.test_recipe_readme \
  tests.test_config.test_example_draft_config_wiring tests.test_runtime.test_package_architecture \
  tests.test_runtime.test_cli_config_build
# Ran 96 tests - OK (skipped=3)

Also green: tests.test_modeling.test_dflash2, test_draft_registry, tests.test_runtime.test_export, test_model_loading. black, isort and autoflake (pre-commit pins) are clean on the changed files.

Not done yet / follow-ups

  • No GPU run. The next step is the DSV4-Flash EP=4 arm on one 8x B200 node (two TP2 capture servers + the DP4 trainer, per the runbook): moe/* load metrics, step time against the frozen-experts arm, and acceptance. The grouped_mm path under EP only runs on CUDA and is untested here.
  • EP adds one host sync per MoE layer per micro-batch (this rank's slot bounds in the sorted routing); test_dflash_host_syncs will need an EP exception once EP recipes join that gate.
  • Partial outputs are reduce-scattered in fp32; a bf16 variant would halve that traffic.
  • Tokens per rank must match across the EP group (true for the DFlash family's fixed anchor shapes).
  • EP combined with moe_freeze_experts (frozen, sliced experts ignored by FSDP) is wired but untested.
  • torchtitan's fixed-capacity / DeepEP dispatchers could be plugged in behind the same gather/scatter seam later; this PR stays on torch 2.13 and adds no dependency.
  • This takes a different route from feat: expert parallelism for draft models with sharded experts #882/feat(moe): partition the routed experts across an expert-parallel group #883 (FSDP and EP together, DTensor slices, token all-gather instead of replicated tokens). Happy to converge with @fg11991 and @sherlockwu on one of them.

yushengsu-thu and others added 12 commits October 2, 2026 09:11
… seams)

Lays out specforge/modeling/draft/moe as one configurable MoE layer with
swappable components, and wires every seam a target-family implementation
needs — with no routing math yet:

- config.py: MoEConfig, moe_preset registry, resolve_moe_config(). Architecture
  keys use the target checkpoints' native HF names at the draft-JSON top level;
  training-only knobs live under dflash_config as moe_*.
- router/balance/experts/shared.py: component contracts + named registries
  (score functions, routers, balance controllers, experts backends, shared
  experts). BalanceController pins the deferred-update timing that keeps
  activation checkpointing correct.
- layer.py: MoELayer composes gate/experts/shared_experts (official attribute
  names); build_ffn() is the dense/MoE switch — dense drafts get the kernel
  provider's MLP verbatim.
- hooks.py: apply_pending_balance_updates / collect_moe_aux_loss /
  collect_moe_metrics over any module tree.
- state_dict.py: to/from_checkpoint_state_dict boundary (module layout vs
  official file naming), applied in the FSDP backend, warm start,
  materialize_draft and both exporters. FSDP full-state-dict hooks need the
  module's own FQNs, so the rename cannot live in state_dict().
- init.py: WarmStartPlan / select_target_experts.
- dflash.py: decoder layers build their FFN through build_ffn; _init_weights
  reaches MoE bare Parameters; the model forward applies pending balance
  updates in training. DFlash/DSpark strategies add moe/* load metrics.
- DESIGN.md + customization docs; tests pin the contracts with stub components.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
RoutingResult carries the differentiable pre-selection scores (optional), and
BalanceController.observe() receives the whole result instead of bare counts,
so a policy can build an auxiliary balance loss for the same forward.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
…ation arm

Fills the MoE skeleton with the DeepSeek-V4 recipe and ships the MoE arm of the
dense-vs-MoE DSpark ablation for DeepSeek-V4-Flash.

- topk_router.py: "topk" router; softmax / sigmoid / sqrtsoftplus score
  functions; optional group-limited selection (DeepSeek top-2 group scores).
  Routing math in fp32; combine weights renormalized then scaled.
- noaux_tc.py: aux-loss-free balancing — fp32 selection bias (kept fp32 through
  module dtype casts) moved by a sign controller on all-reduced expert loads,
  applied from the model forward outside checkpoint regions. Stored in
  checkpoints as gate.bias (the DeepSeek native key SGLang maps onto
  e_score_correction_bias) via a state-dict converter.
- grouped_experts.py: experts as stacked [E, out, in] w1/w2/w3 (FSDP tracks 3
  tensors, grouped GEMMs read them directly); sorted-segment dispatch with a
  torch._grouped_mm path (no host sync) or the portable per-expert loop
  (dflash_config.moe_dispatch). Converter keeps files in the official
  experts.{i}.w{1,2,3}.weight naming. Experts and gate init with the draft's
  initializer_range, like the dense MLP they replace.
- swiglu_shared.py: one ungated SwiGLU shared expert with the V4 clamp.
- presets.py: "deepseek_v4" = sqrtsoftplus + noaux_tc + renorm x1.5 + one
  shared expert + swiglu_limit 10.
- init.py: apply_warm_start() seeds a layer from a target layer's dequantized
  tensors (selected experts, gate rows/bias, shared expert) through the
  checkpoint-naming boundary. Reading/dequantizing DSV4-Flash's fp4 experts is
  left to the target tooling.
- configs/deepseek-v4-flash-dspark-moe.json: the dense DSpark config plus
  moe_preset deepseek_v4, 64 routed + 1 shared experts, top-6, width 2048
  (activated FFN width ~= the dense 12288). The disaggregated recipe mirrors
  the dense one exactly except the draft JSON and run/store names.
- Tests: recipe/config resolution, router weights and group-limited routing,
  deferred bias updates and fp32 buffer, dense per-token reference forward,
  grouped-GEMM vs loop parity (CUDA), official-naming round trips at layer and
  model level, warm start, DFlash training/balancing integration.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
- MoEConfig.serving_fields(): the resolved recipe in the DeepSeek HF config
  vocabulary (scoring_func, topk_method=noaux_tc, routed_scaling_factor,
  n_group/topk_group, norm_topk_prob, swiglu_limit when enabled, ...).
  export --to hf writes them into config.json so a serving engine needs no
  knowledge of SpecForge presets.
- AutoDraftModel.from_pretrained: HF assigns tensors by key and cannot regroup
  per-expert files into the stacked module layout (and refuses an explicit
  state_dict next to a path), so for MoE configs read the safetensors, convert
  through from_checkpoint_state_dict, and load into a freshly built module.
  Warm starts from HF export dirs go through this path.
- Tests: serving fields; an end-to-end export -> config.json -> from_pretrained
  round trip on a tiny MoE DFlash draft.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
dflash_config.moe_aux_loss_coeff > 0 adds DeepSeek-V3's complementary
sequence-wise balance loss (coeff * sum_e f_e P_e over the micro-batch) built
from the router's differentiable scores; DFlash/DSpark strategies add it to
the training loss and it appears as moe/aux_loss.

Why: the bias controller only reorders selection. A from-scratch DFlash-family
drafter feeds the router near-identical mask-token embeddings early on, so
every token routed to the same experts (load_max_ratio pinned at E/k, ~85% of
experts unused through step 100) while the gate logits grew faster than the
bias could follow. The differentiable term bounds the logit gaps so
input-driven routing can emerge.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
- moe/deepseek_v4_target.py: dequantize one target layer's ffn.* tensors
  (packed fp4 e2m1 experts with per-32 ue8m0 scales, fp8 e4m3 shared expert
  with 128x128-block scales, bf16 gate, fp32 noaux bias) into the official
  naming apply_warm_start() consumes; hash-routed layers are rejected.
- scripts/warm_start_moe_drafter.py: build the draft from its config, seed each
  MoE layer from a chosen target layer (identity expert mapping when the
  shapes match, strided subset otherwise), record provenance and the serving
  fields, and write an HF dir usable as model.draft_checkpoint_path.
- Tests for the fp4/fp8 conventions and the layer dequant/naming.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
- AutoDraftModel.from_pretrained (MoE path): honor output_loading_info, which
  the warm-start loader passes; read safetensors memory-mapped and drop the
  regrouped state once loaded.
- warm start: for drafts with a native module layout, read and regroup the
  files directly instead of materializing a second full model per rank, and
  let ranks take turns loading. Eight ranks each building a 65 GB CPU model
  and then loading another 65 GB state + model exceeded the node's RAM and
  got the trainer OOM-killed on the 256-expert warm-started run.
- Test: warm_start_draft_model round-trips a tiny MoE export.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
…FSDP

dflash_config.moe_freeze_experts keeps the routed experts fixed (router,
shared expert and the rest of the draft still train). RoutedExperts opts into
FSDP replication when fully frozen, and the backend's frozen-module scan now
honors that opt-in next to lm_head/embed_tokens.

Why: with DeepSeek-V4-Flash's 256x2048 experts (32.7B params) sharded, every
micro-batch re-gathered ~65 GB of weights and reduce-scattered their grads
(~3 TB/step, 46 s/step on 8 B200s), and at the drafter's LR (6e-4) AdamW
would have overwritten the warm-started experts within ~100 steps anyway.
Frozen + replicated: no expert communication, dense-like step time, and the
pretrained experts are what gets served.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
… stack

Conflict in specforge/training/backend.py: keep the FSDP2 branch's
_sharded_model_state_dict seam and apply the MoE checkpoint-naming
converters (to/from_checkpoint_state_dict) around it, as the MoE stack
does for the FSDP1 path.
…SDP2 backend

Add training.expert_parallel_size (FSDP2 only). Each MoE layer's routed
experts are sliced across that many consecutive ranks as DTensor Shard(0)
over an `ep` mesh axis; FSDP2 shards the slice again over the `efsdp` ranks
through a per-parameter shard_placement_fn while dense parameters stay on
the full data-parallel mesh. The EP axis is carved out of data parallelism,
so ranks keep distinct micro-batches: inside the layer the group all-gathers
its tokens, every rank routes them with the replicated router and computes
only the experts it owns, and a reduce-scatter returns each rank its own
tokens summed over all owners (fixed-size collectives, one host sync per
layer for this rank's slot bounds).

Because every parameter element lives on exactly one rank, the optimizer's
local-shard grad norm, per-shard Adam state and the DCP full-state-dict
checkpoint path are unchanged; the full [E, ...] tensors keep the official
per-expert naming. FSDP2 averages expert gradients over efsdp only while
the owner already summed its group's tokens, so the backend rescales them
by 1/ep at the optimizer boundary. A rank whose experts receive no token
keeps the gathered input, combine weights and expert parameters on the
graph through zero-valued terms so every rank issues the same collectives
and FSDP2 sees the same gradient set.

init_distributed builds the (efsdp, ep) mesh; ParallelConfig carries it;
the FSDP1 backend and NO_SHARD reject EP. Adds the DeepSeek-V4-Flash
DSpark MoE EP=4 recipe and docs.

Tests (CPU gloo): 2-rank layer parity incl. a starved rank, 4-rank
FSDP2TrainingBackend with ep=2 and ep=4 against a data-parallel reference,
checkpoint naming and reload; schema validation.

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

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants