Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
30 commits
Select commit Hold shift + click to select a range
7218a14
feat(distributed): add SM90 FP8 mega MoE
zyy3077 Aug 11, 2026
dc38f23
perf(distributed): use device barriers for mega MoE
zyy3077 Aug 11, 2026
3f4136d
fix(distributed): keep mega MoE barrier offset in range
zyy3077 Aug 11, 2026
c4928cd
perf(distributed): streamline mega MoE dispatch
zyy3077 Aug 11, 2026
e50c55d
perf(distributed): fuse mega MoE execution phases
zyy3077 Aug 12, 2026
8d08b0e
perf(distributed): overlap mega MoE dispatch and compute
zyy3077 Aug 12, 2026
bf81954
perf(distributed): accelerate mega MoE destination dispatch
zyy3077 Aug 12, 2026
fa2f081
perf(distributed): enable fast math for mega MoE L1
zyy3077 Aug 12, 2026
2c7dc26
perf(distributed): vectorize mega MoE output scatter
zyy3077 Aug 12, 2026
2c25e36
perf(distributed): parallelize mega MoE routing
zyy3077 Aug 12, 2026
b56d0d9
feat(distributed): fuse Pro mega MoE path
zyy3077 Aug 12, 2026
3fa62ba
test(distributed): select fused mega MoE path
zyy3077 Aug 13, 2026
23369cf
refactor(distributed): remove mega MoE fallback kernels
zyy3077 Aug 13, 2026
1c6c03d
perf(distributed): broadcast mega MoE scatter metadata
zyy3077 Aug 13, 2026
888c36f
feat(distributed): generalize SM90 mega MoE shapes
zyy3077 Aug 13, 2026
ef751f5
refactor(distributed): clarify mega MoE warp roles
zyy3077 Aug 13, 2026
c3805f4
feat(distributed): schedule mega MoE expert waves
zyy3077 Aug 13, 2026
4be9761
feat(distributed): trace mega MoE kernel pipeline
zyy3077 Aug 14, 2026
00fe052
refactor(distributed): streamline two-kernel mega MoE
zyy3077 Aug 17, 2026
917cc95
Optimize SM90 MegaMoE two-kernel path
zyy3077 Aug 19, 2026
add8350
perf(distributed): retune SM90 mega MoE L1 schedule
zyy3077 Aug 27, 2026
68088ed
perf(distributed): stage mega MoE activation scales in shared memory
zyy3077 Aug 27, 2026
4068861
perf(distributed): make the mega MoE scale pool scale-group major
zyy3077 Aug 27, 2026
c8e2f78
perf(distributed): apply the L1 scale treatment to L2
zyy3077 Aug 27, 2026
a86c8c2
Revert the mega MoE scale-pool layout changes
zyy3077 Aug 27, 2026
ec76e0d
perf(distributed): recalibrate the L2 scatter crossover
zyy3077 Aug 27, 2026
daa40bc
feat(distributed): finalize SM90 FP8 MegaMoE example
zyy3077 Aug 28, 2026
3f0eaa0
refactor(distributed): localize MegaMoE helpers
zyy3077 Aug 28, 2026
2eb007a
refactor(distributed): remove single-kernel MegaMoE prototype
zyy3077 Aug 28, 2026
2e8a973
style(distributed): format SM90 MegaMoE example
zyy3077 Aug 28, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
92 changes: 92 additions & 0 deletions examples/distributed/mega_moe/README.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,92 @@
# SM90 FP8 MegaMoE

This example implements distributed FP8 MegaMoE with two persistent TileScale
kernels on NVIDIA SM90 GPUs:

```text
inputs -> dispatch + L1 GEMM + SwiGLU -> L2 GEMM + scatter + reduce -> output
```

Let:

- `M`: tokens per rank;
- `H`: hidden size;
- `I`: intermediate hidden size;
- `E`: global expert count;
- `R`: rank count;
- `K`: experts selected per token; and
- `C`: per-expert capacity.

Experts are sharded evenly, so each rank owns `E / R` experts.

## Inputs and Output

The pipeline receives the following tensors on each rank:

| Tensor | Shape | Dtype | Description |
| --- | --- | --- | --- |
| `x` | `[M, H]` | FP8 E4M3 | Local input tokens |
| `x_sf` | `[M, H / 128]` | FP32 | Per-128 activation scales |
| `topk_idx` | `[M, K]` | INT32 | Global expert IDs |
| `topk_weights` | `[M, K]` | FP32 | Route weights |
| `l1_weight` | `[E / R, 2I, H]` | FP8 E4M3 | Local gate/up weights |
| `l1_weight_sf` | `[E / R, 2I / 128, H / 128]` | FP32 | L1 per-128 weight scales |
| `l2_weight` | `[E / R, H, I]` | FP8 E4M3 | Local down-projection weights |
| `l2_weight_sf` | `[E / R, H / 128, I / 128]` | FP32 | L2 per-128 weight scales |

The final output is:

| Tensor | Shape | Dtype | Description |
| --- | --- | --- | --- |
| `out` | `[M, H]` | BF16 | Sum of the `K` routed expert outputs for each local token |

## Kernel Boundary

Kernel 1, `fused_l1_swiglu_manual_warp_kernel`, dispatches tokens to their
expert-owning ranks, computes the gate/up projections and SwiGLU, applies route
weights, and requantizes the intermediate activations. The outputs consumed by kernel 2 are:

| Tensor | Shape | Dtype |
| --- | --- | --- |
| `l2_x` | `[E / R, C, I]` | FP8 E4M3 |
| `l2_x_sf` | `[E / R, C, I / 128]` | FP32 |
| `recv_counts` | `[E / R]` | INT32 |
| `src_ranks` | `[E / R, C]` | INT32 |
| `src_tokens` | `[E / R, C]` | INT32 |
| `src_topk` | `[E / R, C]` | INT32 |

Kernel 2, `fused_l2_scatter_reduce_manual_warp_kernel`, consumes these tensors
and the local L2 weights, scatters each routed result back to its source rank,
and reduces the `K` results into `out`. The `combine[M, K, H]` BF16 tensor
is an internal reduction workspace.

## Run

The distributed runtime requires peer-accessible SM90 GPUs and a configured
NVIDIA IMEX channel.

Run a four-GPU correctness smoke test:

```bash
CUDA_VISIBLE_DEVICES=0,1,2,3 \
python examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py \
--num-processes 4 \
--model-config smoke \
--num-tokens 32 \
--capacity 64 \
--check \
--rep 0
```

Benchmark the Flash configuration on four GPUs:

```bash
CUDA_VISIBLE_DEVICES=0,1,2,3 \
python examples/distributed/mega_moe/example_sm90_fp8_mega_moe.py \
--num-processes 4 \
--model-config flash \
--num-tokens 128 \
--capacity 64 \
--warmup 10 \
--rep 100
```
Loading