Add an opt-in fused weighted restore for AutoEP - #8326
Conversation
The AutoEP path moves routed rows around the two expert all-to-alls with
general-purpose tensor ops. It materializes a padded copy of the token matrix,
gathers through an advanced index, scatters expert outputs into a zero-filled
buffer, and builds a [tokens, top_k, hidden] FP32 intermediate only to apply
routing weights and reduce over top-k. Every one of those steps costs a full
pass over the routed activations, and they repeat in every MoE layer.
Add expert_parallel.local_token_backend. It defaults to "eager", which leaves
the current implementation untouched. Setting it to "fused" swaps the reorder
and the weighted restore for Triton kernels that touch each row once.
Writing the reorder out as four passes shows that all of them are the same row
gather, where perm maps an expert-major slot to its source row, inv is its
inverse, and a negative index reads as zero:
permute forward out[i] = tokens[perm[i]] gather by perm
permute backward dtok[j] = dout[inv[j]] gather by inv
unpermute forward out[j] = expert_out[inv[j]] gather by inv
unpermute backward dexp[i] = dout[perm[i]] gather by perm
So one kernel serves all four. perm is injective on real rows, so no pass needs
atomics, and carrying inv costs 4 bytes per row against the 2 * hidden bytes per
row that the padded copy and the zero-filled scatter buffer cost today.
The restore goes straight from [tokens * top_k, hidden] to [batch, seq, hidden],
weighting each row and reducing over top-k in registers. It keeps the eager
dtype discipline: FP32 product and accumulation, one cast on the way out. Its
backward produces both the expert-output gradient and the routing-score
gradient, reducing the score gradient over the hidden dimension per token so
that no cross-program atomic is needed.
The collectives, the router and the grouped GEMM are unchanged, so a measured
difference between the two backends belongs to the local token engine alone.
"fused" is rejected rather than quietly ignored wherever it would have nothing
to replace or would change semantics: folded tensor parallelism, which restores
combined tokens from assignment metadata; an explicit combine_impl="legacy_bmm";
a resolved score_apply other than "post"; and non-CUDA, non-Triton or
non-bf16/fp16 execution, which is checked once before any collective so that
ranks fail together instead of stalling. A run that asked for "fused" and
silently got "eager" would otherwise report the difference between a backend and
itself.
The alignment and index generation that both backends share moves to
generate_local_expert_permute_indices, so the two cannot drift apart.
Signed-off-by: yh0903 <helloyu0903@gmail.com>
The kernel tests cover the reorder and the weighted restore in isolation. They cannot show that a step taken through the fused backend trains the same model, which is the property that decides whether the backend is safe to select. Add a parity test that runs one step through each backend from the same initial state and the same batch, then compares the loss, the block output, the input gradient, every trainable gradient by name, and the parameter delta the optimizer produced. Router and expert gradients reach the comparison by different routes through the restore, so they are asserted to be present rather than left to a bulk comparison that would pass on an empty set. The benchmarked configuration recomputes each MoE block in backward, so the expert-parallel case runs with and without activation checkpointing. A local case with autoep_size=1 covers the branch that skips the all-to-alls. Also fix the fail-fast assertion in the kernel tests: it matched on the required score_apply rather than the rejected one, so it failed against a correct message. Validated on an H100 node: 139 passed, none skipped, for the kernel and config tests, and all three parity cases passed. Signed-off-by: yh0903 <helloyu0903@gmail.com>
Timing the two fused ops against the eager ones they replaced, on an H100 at the
canonical shape, separates them cleanly:
weighted restore forward 2.98x fwd+bwd 1.87x saves 0.215 ms per layer
expert reorder forward 0.89x fwd+bwd 1.01x saves 0.004 ms per layer
The reorder is a wash, and its forward is slower than the eager one. Three
attempts to fix that -- a wider hidden tile, more warps, and several rows per
program -- all landed within noise of the same number. PyTorch's advanced-index
gather and scatter are already close to optimal for this access pattern, and the
fused version additionally has to build an inverse index the eager one does not
need. So it is removed: it carried a kernel, an autograd pair, an index buffer
and a shared-helper refactor, and returned nothing.
An earlier measurement had the reorder at 1.23x. That was an artifact of timing
one iteration at a time and synchronising after each, which charges every
iteration the launch latency a training step hides by queueing work ahead of the
GPU. Timing a batch of iterations between one pair of events measures the
regime these ops actually run in, and the reorder's apparent win disappeared.
What remains is the weighted restore, so it is spelled as what it is: another
implementation of the combine, selected by combine_impl="fused_weighted_sum"
alongside the existing weighted_sum and legacy_bmm, rather than by a second
config key describing a "local token backend" that now moves no tokens. This
also dissolves the question of what a fused backend should do when someone asks
for legacy_bmm: they are alternatives in one enum and cannot both be chosen.
It is still rejected rather than quietly ignored where it has nothing to
replace: folded tensor parallelism, a resolved score_apply other than "post",
and non-CUDA, non-Triton or non-bf16/fp16 execution, checked once before any
collective so ranks fail together instead of stalling.
Validated on an H100 node: 121 kernel and config tests pass, and all three
eager-versus-fused parity cases pass, comparing loss, output, input gradient,
every named gradient, and the optimizer's parameter delta.
Signed-off-by: yh0903 <helloyu0903@gmail.com>
Validate AutoTP folding and expert tensor parallelism independently so neither configuration can mask the other before process-group setup. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: yh0903 <helloyu0903@gmail.com>
Reject malformed tensor contracts before launching Triton, keep malformed permutations deterministic, and use a differentiable PyTorch backward when create_graph requests higher-order autograd. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: yh0903 <helloyu0903@gmail.com>
Keep the optional fused restore unavailable on HIP without importing pytorch-triton-rocm, matching DeepSpeed import-time device compatibility. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: yh0903 <helloyu0903@gmail.com>
Keep comments focused on numerical behavior, fail-fast ordering, and non-obvious kernel choices while removing repeated implementation narration. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: yh0903 <helloyu0903@gmail.com>
There was a problem hiding this comment.
💡 Codex Review
Here are some automated review suggestions for this pull request.
Reviewed commit: c02ea19b5f
ℹ️ About Codex in GitHub
Your team has set up Codex to review pull requests in this repo. Reviews are triggered when you
- Open a pull request for review
- Mark a draft as ready
- Comment "@codex review".
If Codex has suggestions, it will comment; otherwise it will react with 👍.
Codex can also answer questions or update the PR. Try commenting "@codex address that feedback".
Signed-off-by: yh0903 <helloyu0903@gmail.com>
Signed-off-by: yh0903 <helloyu0903@gmail.com>
tohtana
left a comment
There was a problem hiding this comment.
Hi @yh0903,
Thank you for the great improvement on AutoEP. I also reviewed and this PR now looks good to me.
Let's merge after @hwchen2017 also gives the green light.
…ed-weighted-restore # Conflicts: # deepspeed/module_inject/auto_ep_layer.py # deepspeed/module_inject/auto_ep_presets/base.py # docs/code-docs/source/autoep.rst
|
Thanks @hwchen2017 @tohtana for reviewing! I just rebased with master branch and resolved the conflicts. Now CIs are all green, could you help merge this PR when convenient? Thank you! |
…ernels (deepspeedai#8591) # Use int64 program ids in the SwiGLU and AutoEP fused restore Triton kernels Fixes deepspeedai#8590 ## The problem Four Triton kernels compute memory offsets from `tl.program_id`, which is int32. When the tensor has more than 2^31 elements, the offset wraps to a negative number. The mask looks like a bounds check, but a negative offset passes it. The kernels then load and store memory before the start of the tensor, and CUDA reports `an illegal memory access was encountered`. | kernel | int32 product | overflows when | |---|---|---| | `_swiglu_fwd_kernel`, `_swiglu_bwd_kernel` (`swiglu_triton.py`, deepspeedai#8244) | `pid * BLOCK_SIZE`, with `BLOCK_SIZE = 2048` | elements > 2^31 | | `_weighted_restore_forward_kernel` (`autoep_fused_token_ops.py`, deepspeedai#8326) | `token * out_stride`, where the stride is the hidden size | tokens x hidden > 2^31 | | `_weighted_restore_backward_kernel` (same file) | `token * grad_out_stride` | tokens x hidden > 2^31 | For example, in the SwiGLU kernels: ```python pid = tl.program_id(axis=0) # int32 offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) # wraps negative at pid = 1,048,576 mask = offsets < n_elements # a negative offset passes this check ``` ## When it is reached **SwiGLU** runs in `GroupedExperts` on the gate projection, which has shape `[rows received by this rank, moe_intermediate_size]`. With balanced routing, it faults when `tokens per rank x top_k x moe_intermediate_size > 2^31`: | model | top_k | moe_intermediate_size | faults above, tokens per rank | |---|---|---|---| | Qwen3.5-397B-A17B | 10 | 1024 | 209,715 | | DeepSeek-V3 | 8 | 2048 | 131,072 | These rows are arithmetic, not runs. Unbalanced routing reaches the limit sooner, because one rank can receive up to `expert-parallel size x tokens per rank x top_k` rows. We hit it on Qwen3.5-397B-A17B with 8,192-token sequences on 4 nodes x 8 H200. Every token had been routed to one expert-parallel rank of 32, so SwiGLU received 2,621,568 x 1,024 elements (1.25 x 2^31). **The fused restore** is opt-in (`combine_impl="fused_weighted_sum"`). It faults above 524,288 tokens per rank at hidden size 4096, or above 299,593 at hidden size 7168. ## The change Cast the program id to int64 in all four kernels, so every offset product that uses it is computed in int64: ```python # swiglu_triton.py, both kernels pid = tl.program_id(axis=0).to(tl.int64) # autoep_fused_token_ops.py, both _weighted_restore_*_kernel token = tl.program_id(0).to(tl.int64) ``` Each line has a short comment explaining why, because the mask makes the code look safe without it. Two other `program_id` products in `autoep_fused_token_ops.py` are left as they are: - `_invert_index_kernel` multiplies by a block of 256 over `tokens x top_k` indices. - The hidden-dimension tile in the forward kernel is at most the hidden size. ## Testing One H200, torch 2.8.0+cu126, Triton 3.4.0, `CUDA_LAUNCH_BLOCKING=1`, each size in a fresh process. **Existing unit tests:** `tests/unit/v1/ops/triton_ops/test_swiglu_triton.py` and `tests/unit/v1/ops/triton_ops/test_autoep_fused_token_ops.py`, 86 passed. **Above the limit.** Error is the largest absolute difference from a float32 reference: | kernel | size | `master` | this PR | |---|---|---|---| | SwiGLU, bf16 `[rows, 1024]` | 2,149,580,800 elements (1.001 x 2^31) | illegal memory access | forward 0.0132, backward 0.0077 / 0.0077 | | SwiGLU, bf16 `[rows, 1024]` | 2,684,485,632 elements (1.25 x 2^31) | illegal memory access | forward 0.0147, backward 0.0078 / 0.0075 | | fused restore, top_k 2, hidden 4096 | 524,800 tokens (tokens x hidden = 1.001 x 2^31) | illegal memory access | forward 0.0077, row gradient 0.0075, score gradient 7.6e-6 | For comparison, `master`'s SwiGLU errors below the limit are the same size: forward 0.0144 and backward 0.0076 / 0.0075. They come from rounding the float32 result to bf16. **Below the limit, results are bit-identical to `master`.** The sha256 of every output and gradient matches: | kernel | size | tensors compared | |---|---|---| | SwiGLU | 671,612,928 elements (0.31 x 2^31) | output, gate gradient, up gradient | | fused restore | 524,000 tokens (0.9995 x 2^31) | output, row gradient, score gradient | **No time cost.** SwiGLU, median of 50 calls at 671,612,928 elements: | | forward | backward | |---|---|---| | `master` | 0.962 ms | 1.634 ms | | this PR | 0.962 ms | 1.634 ms | Both run at about 4.2 TB/s forward and 4.1 TB/s backward, so memory traffic sets the time. **No new unit test.** A test needs more than 2^31 elements: 12.9 GB for SwiGLU in bf16, and 25.8 GB for the fused restore forward and backward. I can add one with a skip for GPUs that have less free memory, if you want it in CI. ## Reproducer ```python import torch import torch.nn.functional as F from deepspeed.ops.triton_ops.swiglu_triton import swiglu gate = torch.randn(2_099_200, 1024, dtype=torch.bfloat16, device="cuda") # 1.001 x 2**31 elements up = torch.randn(2_099_200, 1024, dtype=torch.bfloat16, device="cuda") out = swiglu(gate, up) torch.cuda.synchronize() # master: "an illegal memory access was encountered" print((out[-1].float() - F.silu(gate[-1].float()) * up[-1].float()).abs().max().item()) ``` On this PR it prints a last-row error of about 0.01. Signed-off-by: pengdurice <pengduhit@gmail.com>
## Summary
- add an opt-in `compile.autoep_non_moe` configuration option;
`engine.compile()` then regionally compiles callable parents of AutoEP
layers
- keep `AutoEPMoELayer.forward` as an explicit compiler-disabled graph
break, so routing, token movement, expert compute, and AllToAll
collectives remain eager
- discover regions from the injected AutoEP module hierarchy, including
a callable model root with a direct AutoEP child
- fail fast for unsupported configurations instead of silently changing
compile behavior
- accept explicitly disabled offload configurations while continuing to
reject active offload
- validate the production eager MoE boundary directly and compare actual
FP32 master-parameter updates
The option defaults to `false`, preserving the existing full-model
behavior of `engine.compile()`. Setting the option alone does not
trigger compilation. Enable it in the DeepSpeed configuration and call
`engine.compile()` after initialization:
```json
{
"compile": {
"autoep_non_moe": true
}
}
```
```python
engine, optimizer, _, _ = deepspeed.initialize(model=model, config=ds_config)
engine.compile()
```
The decoder blocks contain the repeated attention, normalization,
residual and dense work targeted by this optimization. Selecting these
regions bounds the traced code and allows similar blocks to reuse
compiled graphs. Embeddings and the language-model head outside those
regions remain eager; compiling them would need separate validation and
performance measurements. A selected model-root region includes its own
non-MoE operations.
## Support boundary
This first version uses vanilla `torch.compile`/Inductor with the
standard AutoEP `comm` backend, sequence and pipeline parallel sizes of
one, and ZeRO stages 0, 1, and 2. Distributed performance and parity
validation currently target ZeRO stage 1.
DeepEP, DeepCompile, AutoEP+AutoTP folding, sequence or pipeline
parallelism, ZeRO stage 3, optimizer or parameter offload, compiled
autograd, DeepCompile schedules, and any `fullgraph` or `dynamic` value
other than `False` are rejected.
## Historical performance
These measurements predate this review follow-up and the latest master
merge; they are not a new performance measurement of the updated
revision.
Qwen3-30B-A3B, 48 layers, EP16, sequence length 1024, BF16, activation
checkpointing enabled, fixed routing, 2x8 H100. Each allocation
discarded one fixed warm arm and used two runs per variant.
| Allocation | Implementation / order | Eager median | Compiled median |
Throughput gain | Eager p95 median | Compiled p95 median |
| --- | --- | ---: | ---: | ---: | ---: | ---: |
| 1 | pre-production canary, ECCE | 936.66 ms | 885.82 ms | +5.74% |
1627.42 ms | 1534.37 ms |
| 2 | production engine API, CEEC | 932.61 ms | 875.10 ms | +6.57% |
1537.12 ms | 1446.36 ms |
Both allocations reduced peak allocated memory by 1.60 GiB and peak
reserved memory by 2.01 GiB. The maximum paired loss differences were
0.00617 and 0.00194. All 16 ranks reported zero Dynamo counter changes
in the measured window and exactly 2,880 eager AutoEP calls per compiled
arm (`30 steps x 48 layers x forward/replay`).
An exact same-stack Nsight census explained the clean E2E improvement:
| Metric | Eager | Compiled | Delta |
| --- | ---: | ---: | ---: |
| GPU kernels | 22,340 | 15,236 | -7,104 (-31.8%) |
| NCCL kernels | 390 | 390 | unchanged |
| non-NCCL kernel union | 233.73 ms | 191.42 ms | -42.31 ms |
| launch API time | 111.76 ms | 83.07 ms | -28.69 ms |
Compiled execution also removed 1.59 GiB of D2D traffic per captured
step. These performance runs used the PR1+PR2+PR3 benchmark stack
(deepspeedai#8326, deepspeedai#8331, and deepspeedai#8359) to isolate the remaining Transformer/runtime
fragmentation. This PR is based directly on `master` and has no code
dependency on those changes.
## Project rebaseline
This is context for the broader AutoEP optimization effort, not a merge
gate for this PR. A same-allocation fixed-routing `discard + ABBA +
BAAB` comparison of the full deepspeedai#8326 + deepspeedai#8331 + deepspeedai#8359 + deepspeedai#8380 stack against
Megatron Core measured:
| Framework | Median step | Median of per-arm p95 | Peak allocated |
Peak reserved |
| --- | ---: | ---: | ---: | ---: |
| AutoEP PR1-4 | 887.21 ms | 1244.49 ms | 40.37 GiB | 45.29 GiB |
| Megatron Core | 814.58 ms | 821.30 ms | 41.68 GiB | 45.80 GiB |
The residual in that historical comparison was **72.64 ms / AutoEP
+8.9%**. All four paired deltas favored Megatron and stayed within
**66.55–77.19 ms**; AutoEP used 1.30 GiB less allocated and 0.51 GiB
less reserved memory.
A separate low-overhead outer CUDA-event trace placed the typical
residual primarily in:
- forward: **AutoEP +39.14 ms**
- backward including ordinary gradient synchronization: **AutoEP +46.27
ms**
- optimizer: **AutoEP -9.00 ms**
The AutoEP tail is not a measured-window recompilation. In both traced
AutoEP arms, step 27's forward increased from a typical 266–271 ms to
678 ms while backward and optimizer remained near their medians. With
only ten measured steps per arm, the reported p95 is the maximum sample;
this deterministic forward-only spike is a follow-up investigation
rather than evidence that regional compile regressed the typical path.
## Validation
- 57 targeted CPU tests pass: configuration parsing, unchanged
full-model compile behavior, callable-root forward/backward parity with
checkpointing off/on, fail-fast and rollback contracts, parity
assertions, and existing DeepCompile lifecycle tests.
- Both callable-root regression cases reproduce the erroneous rejection
with the original helper and pass after the fix.
- `pre-commit run --files` passes for all seven changed files; all
non-merge commits carry `Signed-off-by` trailers.
- Exact-head GPU validation of
`85ba684c8709bfe1099c41151f504d83d7344636`: all **8/8 cases passed**,
with zero failures, errors or skips, on **2 × NVIDIA H100 80GB HBM3**
(PyTorch **2.10.0.2+cu130**, CUDA **13.0**, NCCL **2.28.9**). Coverage
is nested/root layouts × checkpointing off/on × async split planning
off/on, using BF16 and ZeRO stage 1.
The GPU suite uses the configuration option through
`deepspeed.initialize()` and the unchanged `engine.compile()` call. It
checks actual Inductor graph capture, output/loss/input gradients,
router/expert/non-MoE gradients, FP32 master-parameter updates, exact
routing assignments, the production eager AutoEP boundary, and no
additional Dynamo graphs/calls after warmup. These targeted tests are
correctness checks, not a new performance measurement. The single
measured SGD step uses `lr=1.0` to keep FP32 master updates above
rounding near unit-valued LayerNorm parameters. At `lr=0.01`, the
root-layout update norm was about `4e-7`, so a difference of one FP32
ULP exceeded the 5% relative gate. The forward/gradient comparisons and
every acceptance threshold are unchanged.
The current `modal-torch-latest` workflow runs test selection on PR
updates and executes its GPU suite on merge-queue entries. The
historical full-CI timeout is not counted as a completed test run.
---------
Signed-off-by: yh0903 <helloyu0903@gmail.com>
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Summary
combine_impl="fused_weighted_sum"path while keeping the existing eager weighted reduction as the default.[tokens * top_k, hidden]expert rows directly to[batch, sequence, hidden]in one Triton pass with FP32 weighting and accumulation, including expert-row and routing-score gradients.Performance
[-1.12%, +3.02%]; this is not yet a statistically conclusive end-to-end win.Testing Done
Local code review completed
Unit tests added/updated
Integration tests pass
Manual testing performed
Repository pre-commit hooks passed for all nine changed files.
H100 fused token-op suite: 23 passed, including top-k 2/4/6/8, non-power-of-two hidden sizes, input-contract validation, forward/backward parity, and double backward.
H100 configuration tests: 5 passed for standard, folded AutoTP, expert tensor parallelism, and score-application validation.
Full-step eager/fused parity: 3 passed, covering loss, output, input gradients, router/expert gradients, optimizer parameter deltas, activation checkpointing on/off, EP2, and local experts.
Earlier final restore-only sweep: 121 kernel/config tests passed before the additional review-driven safety cases were added.