[Feature] Decouple expert parallelism from FSDP sharding (dp2ep layout) - #2093
Open
silencelamb wants to merge 14 commits into
Open
silencelamb wants to merge 14 commits into
silencelamb wants to merge 14 commits into
Conversation
…layout Single-process fake ProcessGroup tests asserting the mesh ranks and DTensor placements that MoE.fully_shard produces today for 8/16/64 ranks, plus the Phase-0 baseline snapshot and the decision log for the EP/FSDP decoupling work.
Add FSDPConfig.decouple_ep_fsdp. When enabled, MoE builds a single (replicate, efsdp, ep) root mesh with efsdp = dp_shard / ep, shards routed experts with their own FSDP group on efsdp and every other parameter over the full dp_shard mesh (no more ep-fold replication of dense params), and scales expert gradients by 1/ep instead of all-reducing dense gradients over EP. The legacy layout is untouched when the flag is off. Includes fake-PG L0 placement tests for 8/16/64 ranks, the 8-GPU numerics script used for the L1 report, and reports/L1.md.
… EP/FSDP layout - Float8Handler: per-parameter-class fp8 padding chunk counts and tile-wise reduce meshes (routed experts follow the efsdp shard stride, dense params the flattened dp_shard); legacy single-mesh path unchanged. - BaseModel._fsdp_foreach_allgather picks the FSDP gather group per LoadSpec so EP-local expert slices are still reconstructed for RL weight sync. - examples/v1/config/sft_glm5p2.py: DECOUPLE_EP_FSDP / HSDP_SHARDING_SIZE env switches. - L0 fake-PG tests for the fp8 meshes, the L3 DCP/HF checkpoint script and reports/L3.md (bit-exact HF export, DCP resume, cross-layout DCP reshard, fp8 numerics within the fp8 run-to-run noise floor).
Root cause: `WeightIterator.iter_layer_batches` (IPC + Turbomind) took the parameters and `LoadSpec`s of a compose model's language tower but ran the FSDP-only all-gather through the outer compose model. The compose model is wrapped on the world mesh, has no `expert_fsdp_mesh`, and never builds a `load_spec_mapping` of its own, so on the decoupled EP/FSDP layout the language tower's `efsdp` expert shards (and, with HSDP, its `dp_shard` dense shards) were treated as shards to preserve and never gathered: Turbomind received rank-local fragments. The vision / projector / non-layer language parameters additionally hit the empty `load_spec_mapping` of the compose model. Fix: resolve the owning submodule (language_model / vision_tower / multi_modal_projector, or the model itself) for every parameter and run `_fsdp_foreach_allgather`, `_to_float8` and the `load_spec_mapping` lookup on that owner. Plain models are unaffected. Test plan: `tests/rl/test_weight_iterator.py::TestLayerBatchesGatherWithParamOwner` builds a compose model whose language tower is EP-sharded and FSDP-sharded on `efsdp` (with and without an HSDP replicate dim) on fake process groups and asserts every streamed tensor is complete (EP slice kept; `efsdp`, `dp_shard` and world shards gathered). It fails before the fix and passes after; the existing weight-iterator, load-spec and L0 mesh tests still pass.
…inks, limitation details
- FSDPConfig: replace the `assert`s in `model_post_init` with a `@model_validator(mode="after")` that raises `ValueError` (surfaced as a pydantic `ValidationError`), so the checks survive `python -O`; also reject non-positive `ep_size` / `hsdp_sharding_size`. - MoE._init_decoupled_device_mesh: raise `ValueError` instead of asserting `world_size % dp_shard == 0` and `dp_shard % ep_size == 0`. - MoE._scale_and_reduce_grad_decoupled: narrow `mesh_dim_names` before indexing so the decoupled path type-checks. - L0 regression tests for the config validator and the runtime check.
- xtuner/_testing/decoupled_ep_fsdp.py: the tiny Qwen3-MoE checkpoint, token stream, per-layout training run and HF-checkpoint comparison shared by the gates and the manual experiment scripts. - tests/engine/test_decoupled_ep_fsdp_train_engine.py: 8-GPU `DeterministicDDPTestCase` gates. L1 / L2 compare the decoupled layouts (efsdp == 1, efsdp > 1, HSDP + EP) with the legacy ones on loss curves, total grad norms, per-parameter grad norms at step 0 and per-rank parameter memory; L3 checks bit-exact HF export after `from_hf`, DCP resume, the HF export after resume and cross-layout DCP resharding. CPU tests pin the checkpoint comparison itself. - run_decoupled_ep_fsdp_numerics.py / run_decoupled_ep_fsdp_ckpt.py now import the shared helpers; their CLI and JSON output are unchanged.
…upled EP/FSDP design note - §3.1 shows the `model_validator`; §3.6 and the §5 evidence table link the 8-GPU gates next to the experiment scripts; the two resolved items leave §6.3. - reports/decoupled_ep_fsdp_review_zh.md: status note for N1 and W5.
silencelamb
force-pushed
the
pr/decouple-ep-fsdp
branch
from
September 27, 2026 12:19
8f03bc2 to
ed627b4
Compare
…after every call On the decoupled path each MoEBlock is its own FSDP unit nested inside the layer's reentrant activation checkpoint. With FSDP2's automatic reshard the group was resharded after every call: once per intra-layer micro-batch in forward, again for every call the checkpoint replays in backward and once per micro-batch backward, each followed by a fresh all-gather. With four micro-batches that is 13 expert all-gathers per layer and step where the legacy layout, whose single FSDP unit per layer wraps the checkpoint, needs 2. Expert groups are now sharded with reshard_after_forward=False and reshard_after_backward=False and resharded explicitly: by a forward hook on the outer layer unit (checkpoint replay does not call it) and by a hook on the gradient of the layer inputs, which fires after the whole layer backward. A model forward pre-hook reshards any group left gathered so no step can run on pre-update weights. MTP layers use the same two-level wrap and get the same treatment. The expert groups are all-gathered once in forward and once in backward per layer, and at most one layer's experts are gathered at a time.
…elow FSDP2 picks each unit's default backward prefetch target from the order in which units finished forward. The decoupled expert units sit inside the reentrant activation checkpoint, so checkpoint replay re-runs their forward hooks during backward and appends to that order. Their default target then becomes the experts of the layer above, whose backward is already done: one all-gather per layer that is never used and whose output buffer stays allocated until the end of backward (0.3 GB per layer on Qwen3-30B-A3B at EP4, ~14 GB in total, which set the peak memory of the decoupled layout). Nothing prefetches the dense unit of the layer below either, so its all-gather is issued on demand and queues behind this layer's dense reduce-scatter on the same process group. Every expert group now prefetches the layer below (for decoder layers and MTP layers alike); an explicit list replaces the default target. The first layer of each stack keeps the default.
…ches and recompute The existing gates compare losses, gradient norms and per-rank parameter memory, none of which changes when an FSDP group is gathered more often than needed, and they run without intra-layer micro-batches. The new 8-GPU gate trains the tiny Qwen3-MoE at EP4 (efsdp == 2) with two intra-layer micro-batches and the default full recompute, counts per FSDP unit the all-gathers issued and the ones whose result was used, and asserts at most two used all-gathers per expert group and step (forward + backward) and at most one unused all-gather in total. Before the fix it fails with six used all-gathers per expert group. `count_fsdp_all_gathers` wraps `FSDPParamGroup.unshard` / `wait_for_unshard` to collect the counts; `build_engine` / `run_mode` take the number of intra-layer micro-batches.
…ard prefetch Add §3.3.1 to the design note: why the expert FSDP groups, nested inside the layer's reentrant activation checkpoint, need per-layer reshard and an explicit backward prefetch target, and record the measurements in §5. Refresh the moe.py line references.
silencelamb
force-pushed
the
pr/decouple-ep-fsdp
branch
from
September 27, 2026 22:47
ed627b4 to
7c53b7d
Compare
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.
Summary
This PR decouples MoE expert parallelism from FSDP sharding (the "dp2ep" layout of torchtitan): EP becomes a sub-dimension of the FSDP shard dimension instead of an orthogonal mesh axis. Dense parameters are sharded over the full data-parallel shard group and are no longer replicated across EP ranks; routed experts keep their EP split and are FSDP-sharded only over the remaining
efsdp = dp_shard / epranks. The switch isFSDPConfig.decouple_ep_fsdp(defaultFalse; the legacy layout is untouched). The design note isdocs/design/decouple_ep_fsdp.md.flowchart LR R["root mesh (replicate, efsdp, ep)"] --> EP["ep_mesh = root[ep]<br/>token dispatch / expert ownership"] R --> DS["fsdp_mesh = flatten(efsdp, ep) = dp_shard<br/>FSDP group of dense params"] R --> EF["expert_fsdp_mesh = root[efsdp]<br/>FSDP group of routed experts on top of EP"] R -. "replicate > 1 (HSDP)" .-> HS["hsdp_mesh = root[replicate, dp_shard]<br/>expert_fsdp_mesh = root[replicate, efsdp]"]With the legacy
(fsdp, ep)mesh,fsdp = world / ep, so EP=8 on 8 GPUs keeps a full copy of every dense parameter and its fp32 optimizer state on every rank. For GLM-5.2-30B that pins the allocator at its ceiling (2–3 allocation retries per step) and makes EP8 about 7× slower than EP4. With the decoupled layout EP4 goes from 99.5 to 83.5 GiB at equal step time, and EP8 from 120.0 GiB / 11.8 s to 76.3 GiB / 1.64 s per step (8×H200, production recipe).Changes
FSDPConfig.decouple_ep_fsdp, validated by a pydanticmodel_validator(ep | dp_shard, positive sizes); the runtime mesh checks raiseValueError.MoE._init_decoupled_device_mesh: one root mesh(replicate, efsdp, ep);fsdp_mesh = root[efsdp, ep]._flatten("dp_shard")keeps its meaning for existing consumers,expert_fsdp_meshis new, and HSDP + EP coexist on this path.fully_shard: everyMoEBlockis wrapped onexpert_fsdp_meshfirst, then each decoder / MTP layer on the dense mesh; expert all-gathers are prefetched together with the next layer. The expert groups sit inside the layer's reentrant activation checkpoint, so FSDP2's per-call automatic reshard is off for them: hooks on the outer layer unit reshard them once after the layer's forward and once after its backward, and their backward prefetch points at the layer below._scale_and_reduce_grad_decoupled: routed experts getgrad.div_(ep)on top of FSDP's reduce-scatter, FSDP-ignored fully replicated fp32 parameters get one coalesced all-reduce, and the legacy manual cross-EP all-reduce of dense gradients is gone.dp_shardchunks, rank stride 1; experts:efsdpchunks, rank strideep).LoadSpecderives the shard history from the new DTensor placements.BaseModel._fsdp_gather_grouppicks the gather group perLoadSpec, andWeightIterator._param_ownergathers compose-model parameters with the submodule that owns them, soefsdp/ HSDP shards of a language tower are gathered instead of being mistaken for EP-local shards.examples/v1/config/sft_glm5p2.py:DECOUPLE_EP_FSDPandHSDP_SHARDING_SIZEenvironment switches.tests/model/share their helpers with the gates.Verification
xtuner/v1files (codespell, docformatter, pyupgrade, ruff, ruff-format, pydantic-extra-check) passed; mypy 1.16.1 reports no errors in the lines this PR touches.tests/model/test_decoupled_ep_fsdp_mesh.py: 31 passed (fake process groups, CUDA context only).tests/rl/test_weight_iterator.py: 13 passed.tests/engine/test_decoupled_ep_fsdp_train_engine.pyon 8×H200 (torch 2.9.1, tiny random Qwen3-MoE): 7 passed. L1 (ep=1, ep=8 legacy, ep=8 decoupled) and L2 (ep=4 legacy vs.efsdp=2, HSDP + EP withefsdp=1andefsdp=2): loss curves within 1e-3 relative (observed ~2e-5), total and step-0 per-parameter gradient norms within 1e-2, per-rank parameter memory matching the layout. L3: bit-exact HF export right afterfrom_hf, DCP resume, HF export after resume within bf16 ulps, and cross-layout DCP resharding. The all-gather gate allows at most 2 used all-gathers per expert group and step.Result
decouple_ep_fsdpdefaults toFalseand the execution logic of the legacy path is unchanged; the legacy placements are pinned by the L0 tests. Known limits are listed in §6 of the design note: ExpertTP and decoupling are mutually exclusive,fp32_keys_patternmatching a routed expert is unsupported on both paths, HSDP + EP is validated on single-node scaled-down topologies only, and the two-meshfully_shard+DeviceMesh._flattencombination is validated on torch 2.8 / 2.9.