Skip to content

Add an opt-in fused weighted restore for AutoEP - #8326

Merged
hwchen2017 merged 15 commits into
deepspeedai:masterfrom
yh0903:yh0903/autoep-fused-weighted-restore
Sep 2, 2026
Merged

hwchen2017 merged 15 commits into
deepspeedai:masterfrom
yh0903:yh0903/autoep-fused-weighted-restore

Conversation

@yh0903

@yh0903 yh0903 commented Aug 26, 2026 •

Copy link
Copy Markdown
Contributor

Summary

  • Add an opt-in AutoEP combine_impl="fused_weighted_sum" path while keeping the existing eager weighted reduction as the default.
  • Restore [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.
  • Fail fast for unsupported CUDA/dtype/score/AutoTP/expert-TP configurations, preserve higher-order autograd through a differentiable fallback, and document the experimental path.

Performance

  • Canonical-shape H100 microbenchmark: 2.98x forward and 1.87x forward+backward, saving about 0.215 ms per layer.
  • Qwen3-30B-A3B, 48 layers, EP16, 2x8 H100: pooled median improved by 0.97%, with a 95% CI of [-1.12%, +3.02%]; this is not yet a statistically conclusive end-to-end win.
  • Peak reserved memory decreased by 96 MiB, matching the removed 64 MiB FP32 weighted intermediate and 32 MiB assignment buffer.
  • The experimental expert reorder was removed after launch-amortized measurements showed no benefit, so this PR changes only the weighted restore.

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.

yh0903 and others added 8 commits August 24, 2026 16:56
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>
@yh0903
yh0903 marked this pull request as ready for review August 26, 2026 23:23

@chatgpt-codex-connector chatgpt-codex-connector Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

💡 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".

Comment thread deepspeed/moe/autoep_fused_token_ops.py Outdated
Comment thread deepspeed/moe/autoep_fused_token_ops.py Outdated
Comment thread deepspeed/moe/autoep_fused_token_ops.py Outdated
Comment thread deepspeed/moe/autoep_fused_token_ops.py Outdated
Comment thread tests/unit/v1/moe/test_autoep_fused_token_ops.py Outdated
Comment thread deepspeed/moe/autoep_fused_token_ops.py Outdated
Comment thread deepspeed/ops/triton_ops/autoep_fused_token_ops.py
Comment thread deepspeed/moe/autoep_fused_token_ops.py Outdated
Comment thread deepspeed/moe/autoep_fused_token_ops.py Outdated
Comment thread deepspeed/ops/triton_ops/autoep_fused_token_ops.py Outdated

@tohtana tohtana left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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
@yh0903

yh0903 commented Sep 2, 2026

Copy link
Copy Markdown
Contributor Author

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!

@hwchen2017
hwchen2017 added this pull request to the merge queue Sep 2, 2026
Merged via the queue into deepspeedai:master with commit 80f19f3 Sep 2, 2026
13 checks passed
pull Bot pushed a commit to dumpmemory/DeepSpeed that referenced this pull request Sep 19, 2026
…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>
pull Bot pushed a commit to Abaso007/DeepSpeed that referenced this pull request Sep 22, 2026
## 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>
@tohtana tohtana mentioned this pull request Sep 27, 2026
22 tasks
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.

3 participants