(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
Draft
yushengsu-thu wants to merge 12 commits into
yushengsu-thu wants to merge 12 commits into
Conversation
… 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
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.
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 shardedfsdp_sharding,tp_size=1, no sequence parallelism, a world size divisible by it, andn_routed_expertsdivisible by it).init_distributedbuilds a 2-D(efsdp, ep)mesh (ep= consecutive ranks, one node for ep <= GPUs per node);ParallelConfigcarries it.specforge/modeling/draft/moe/expert_parallel.py: the EP layout,shard_expert_parameter(the stacked[E, ...]weights becomeDTensorShard(0)overep), and the differentiable seamsgather_tokens(all-gather, backward reduce-scatter) andscatter_outputs(reduce-scatter, backward all-gather).MoELayer.forwardunder 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_paralleland a local-slice dispatch for bothsorted_loopandgrouped_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-parametershard_placement_fn(ShardPlacementResult, available in torch 2.13) puts expert slices on theefsdpmesh and everything else on the full mesh; rescales expert gradients by1/epat 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.NO_SHARDreject EP with explicit errors.examples/configs/online/disaggregated/external/deepseek-v4-flash-dspark-moe-ep4-disaggregated.yaml(the MoE arm withbackend: 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 overefsdp = world / ep, but an expert owner's gradient already sums its whole group's tokens, so the backend multiplies expert gradients by1/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-rankFSDP2TrainingBackendwith 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 throughset_model_state_dict; the unsharded layer is unchanged.tests/test_config/test_schema.py: the new field's validators and world-size check.Also green:
tests.test_modeling.test_dflash2,test_draft_registry,tests.test_runtime.test_export,test_model_loading.black,isortandautoflake(pre-commit pins) are clean on the changed files.Not done yet / follow-ups
moe/*load metrics, step time against the frozen-experts arm, and acceptance. Thegrouped_mmpath under EP only runs on CUDA and is untested here.test_dflash_host_syncswill need an EP exception once EP recipes join that gate.moe_freeze_experts(frozen, sliced experts ignored by FSDP) is wired but untested.