exp.[rebase] MoE FFN for DFlash-family drafters (#811/#812) on current main, to run the DeepSeek-V4-Flash MoE arm - #931
yushengsu-thu wants to merge 8 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>
Smoke run of the DeepSeek-V4-Flash DSpark MoE arm on this branch (GB300, 50 optimizer steps)Setup. One 4x GB300 node (284 GB each, aarch64, Three 50-step runs with identical overrides, capture servers kept up and only
What works.
What to watch.
Logs: 🤖 Generated with Claude Code |
What the rebase touched
kan/moe-2-dsv4was 151 commits behindmain.git rebase origin/mainreplayed the eight commits cleanly except one hunk indocs/recipes/deepseek-v4-flash-dspark-disaggregated.md, wheremainhad added the "On AMD MI355X" section at the same place the stack adds "MoE-FFN arm (ablation)". Both sections are kept, AMD first.mainhad already moveddocs/advanced_features/customization.mdtodocs/sections/advanced_features/; git followed the rename, so the stack's MoE section landed in the new location.Checks
CPU, torch 2.13.0 / transformers 5.12.1, on the rebased tree:
The
lintjob will inherit the pre-existing autoflake findings on the stack's own files that @fg11991 noted on #883; they are not introduced by the rebase.Plan for this draft
Run
examples/configs/online/disaggregated/external/deepseek-v4-flash-dspark-moe-disaggregated.yaml(two TP2 capture servers + DP4 trainer on one 8x B200 node, per the runbook) on this branch and attach the run notes here: whether the recipe starts and trains on currentmain,moe/*load metrics, step time, and any fixes the rebase needs.🤖 Generated with Claude Code