Skip to content

perf(qwen3): Qwen3-30B-A3B FP8 on MI355X to 40k tokens/s/GPU with Primus-Turbo tensorwise - #1247

Draft
Xiaoming-AMD wants to merge 6 commits into
mainfrom
perf/qwen3-30b-a3b-fp8-turbo
Draft

Xiaoming-AMD wants to merge 6 commits into
mainfrom
perf/qwen3-30b-a3b-fp8-turbo

Conversation

@Xiaoming-AMD

Copy link
Copy Markdown
Collaborator

Description

Qwen3-30B-A3B FP8 pretraining on 8× MI355X goes from a memory fault at iteration 3 to 40,283 tokens/s/GPU (6507.6 ms/iter), 47% above the BF16 config. This PR switches the FP8 example config to Primus-Turbo tensorwise FP8 and turns on every Turbo fusion that got there. Two of those fusions are new in this PR:

  • head_dim 128 for the fused QKV split + Q/K RMSNorm + RoPE patch. The patch was GPT-OSS-only before.
  • turbo_fp8_permute: the DeepEP permute hands the fused grouped MLP FP8 tokens, and takes its gradient back in FP8.

The PR also fixes enable_turbo_attention_float8, which failed on its first call.

Primus-Turbo requirements.

  • turbo_fp8_permute needs a Primus-Turbo whose moe_permute takes quantize_dtype (AMD-AGI/Primus-Turbo#555). With an older Turbo, the first dispatch fails.
  • The head_dim 128 fusion needs AMD-AGI/Primus-Turbo#556. Without it, the patch keeps the unfused path.
  • The other Turbo rows in the table below come from #554, #557 and #558. They need no Primus change.

End-to-end. Qwen3-30B-A3B FP8 pretrain on 8× MI355X (EP8, MBS 8, GBS 512, seq 4096, even routing), mean of iterations 11–20 of 20-iteration runs. Each row adds one change on top of the row above:

run from ms/iter TFLOP/s/GPU tokens/s/GPU mem loss@20
BF16 config (legacy grouped GEMM) 9587 622.1 27,344
FP8 config before this PR memory fault at iteration 3
FP8 tensorwise config (Turbo e2d9f1d7) this PR 8894.0 670.6 29,474 277 GB 11.33393
+ use_turbo_fused_act_with_probs this PR 7915.6 753.5 33,117 223 GB 11.33435
Turbo 1103b2df Turbo main 7925.4 752.5 33,077 222 GB 11.33427
+ flat tensorwise FP8 quant kernel Turbo #558 7870.9 757.7 33,305 222 GB 11.33415
+ per-token get_dispatch_layout Turbo #557 7774.2 767.2 33,720 222 GB 11.33431
+ HIP permute in DeepEPTokenDispatcher Turbo #554 7643.5 780.3 34,296 222 GB 11.33423
+ turbo_deepep_num_cu 80 → 160 this PR 7353.4 811.1 35,649 223 GB 11.33456
+ fused QKV split + Q/K RMSNorm + RoPE, head_dim 128 this PR + Turbo #556 6987.1 853.6 37,518 224 GB 11.33373
+ turbo_fused_grouped_gemm (mean of 2 runs) this PR 6817.7 874.9 38,451 225 GB 11.33412
+ turbo_fp8_permute (mean of 4 runs) this PR + Turbo #555 6507.6 916.7 40,283 223 GB 11.33404

From the FP8 tensorwise config to the last row: −26.8% ms/iter, +36.7% tokens/s/GPU. Loss differences are within run-to-run noise (about 4e-4 at iteration 20). This branch, rebased on main 0a68e5c, with the switches on:

run ms/iter tokens/s/GPU loss@20
switches on the command line 6518.7 40,214 11.33415
this config alone 6410.6 40,892 11.33433
this config alone 6459.6 40,582 11.33445

Known issue: intermittent NaN.

  • One more config-only run of this branch hit a NaN forward loss on rank 4 at iteration 3. Its parsed arguments match a passing run exactly.
  • Across this work, 1 of 13 20-iteration runs at turbo_deepep_num_cu 160 hit this, as did the only run at 192 (rank 2, iteration 5).
  • None of the 5 runs at 80 CUs did. Those runs also predate the fused paths, though, so the cause is not isolated yet.

Changes

  • fix(turbo): make enable_turbo_attention_float8 runnable again
    • primus/backends/megatron/core/extensions/primus_turbo.py:
      • Pass sink= only when a sink tensor exists; flash_attn_fp8_func has no sink parameter.
      • Make q/k/v bshd-contiguous on the FP8 path. flash_attn_fp8_func permutes sbhd storage as if it were the logical layout (fixed in Turbo by fix/fp8-attn-strided-layout).
  • perf(qwen3): tune Qwen3-30B-A3B FP8 MI355X config to Turbo tensorwise
    • examples/megatron/configs/MI355X/qwen3_30B_A3B-FP8-pretrain.yaml: fp8_recipe: tensorwise, use_turbo_gemm / use_turbo_grouped_gemm, and the TE grouped-MLP path instead of the legacy grouped GEMM.
  • perf(megatron): enable the fused packed-QKV RMSNorm + RoPE patch for head_dim 128
    • primus/backends/megatron/patches/turbo/qk_rmsnorm_rope_patches.py:
      • Accept any head_dim in Turbo's QK_RMSNORM_ROPE_HEAD_DIMS, falling back to (64,) on older Turbo.
      • Accept TE RMSNorm as well as PrimusTurboRMSNorm. The fused path only reads weight / eps; zero_centered_gamma is still rejected.
    • Fused QKV split + Q/K RMSNorm + RoPE at the Qwen3 attention shape: fwd+bwd 1243.6 → 369.9 µs (3.36×).
  • perf(megatron): hand the Turbo fused grouped MLP FP8 tokens from the DeepEP permute
    • New opt-in flag turbo_fp8_permute (default false in primus_turbo.yaml).
    • Under Turbo FP8 tensorwise current scaling, PrimusTurboDeepEPTokenDispatcher passes quantize_dtype to _post_dispatch and grad_quantize_dtype to _pre_combine. The dtypes follow the FP8 format (HYBRID: e4m3 input, e5m2 gradient), and results are bit-identical to the default path. Any other recipe, or a layer with Turbo FP8 off, keeps the bf16 path.
    • Argument validation requires enable_primus_turbo, use_turbo_deepep and turbo_fused_grouped_gemm, and rejects selective recompute of moe_act.
    • Dispatcher kernel time per call at the Qwen3 EP8 shape: forward 645 → 167 µs, backward 687 → 174 µs.
  • perf(qwen3): turn on the Turbo fusions in the Qwen3-30B-A3B FP8 MI355X config
    • Top-level env: PRIMUS_FUSED_QK_RMSNORM_ROPE: "1".
    • turbo_fused_grouped_gemm, use_turbo_fused_act_with_probs and turbo_fp8_permute set to true.
    • turbo_deepep_num_cu 80 → 160. 192 gave a NaN loss on one rank at iteration 5; 256 is slower (7028 ms).

The force-"even" routing fix that this work also needed is already on main (#1243), so it is not part of this PR.

Tests

  • New in tests/unit_tests/backends/megatron/test_rocm_arg_validation.py: validate_turbo_fp8_permute (each required flag, selective recompute of moe_act, disabled no-op).
  • New tests/unit_tests/backends/megatron/test_turbo_fp8_permute_dtypes.py: _fp8_permute_dtypes returns (e4m3, e5m2) for HYBRID and (e4m3, e4m3) for E4M3 tensorwise, and (None, None) with the flag off, Turbo FP8 off or blockwise scaling.
  • test_qk_rmsnorm_rope_patches.py, test_rocm_arg_validation.py, test_validate_args_patches.py, test_router_force_even_routing.py, test_turbo_fp8_permute_dtypes.py: all pass on MI355X.
  • End-to-end runs above.

xmpeng-dev and others added 5 commits October 10, 2026 15:17
Turbo FP8 attention failed on its first call for two reasons:

- The core-attention call always passed ``sink=`` (added for sink attention),
  but ``flash_attn_fp8_func`` has no ``sink`` parameter, so every FP8 run
  raised a TypeError. Pass ``sink`` only when a sink tensor exists.
- Megatron hands over sbhd storage viewed as bshd. ``flash_attn_fp8_func``
  infers "sbhd" from the strides and then permutes the already-bshd logical
  tensor, swapping b and s (``shape '[4096, 32, 0, 64, 128]' is invalid``).
  Make q/k/v bshd-contiguous on the FP8 path until the Turbo fix lands.
Switch the FP8 recipe to tensorwise with dense and grouped GEMMs on
Primus-Turbo (FlyDSL grouped GEMM on gfx950) and the TE grouped-MLP path
instead of the legacy grouped GEMM. The previous config hit a GPU memory
fault (sum_and_scatter) at iteration 3 on the current image.

1 node x 8 MI355X, EP8, MBS 8, GBS 512, seq 4096, force-even routing:

| config                         | ms/iter | TFLOP/s/GPU | tokens/s/GPU |
| ------------------------------ | ------- | ----------- | ------------ |
| BF16 config (legacy grouped)   |    9587 |       622.1 |       27,344 |
| FP8 config, before             |   memory fault at iter 3           |
| FP8 config, this change        |    8894 |       670.6 |       29,474 |
…head_dim 128

The PRIMUS_FUSED_QK_RMSNORM_ROPE=1 patch only fired for GPT-OSS: it required
head_dim 64 and PrimusTurboRMSNorm Q/K norms. Qwen3-30B-A3B uses head_dim 128
and, with use_turbo_rms_norm false, TE RMSNorm, so it always fell back.

- Accept any head_dim listed in Turbo's QK_RMSNORM_ROPE_HEAD_DIMS (64 and 128
  with the companion Turbo change; falls back to (64,) on older Turbo).
- Accept TE RMSNorm; PrimusTurboRMSNorm subclasses it and the fused path only
  reads weight / eps (zero_centered_gamma is still rejected).

Qwen3-30B-A3B attention shape on MI355X (S=4096, B=8, 4 KV groups x 8 Q heads
per group, D=128): split + Q/K RMSNorm + RoPE fwd+bwd 1243.6 -> 369.9 us
(3.36x); 48 layers x 8 micro-batches is about -335 ms per iteration.

End-to-end: Qwen3-30B-A3B FP8 tensorwise pretrain on 8x MI355X (EP8, MBS 8,
GBS 512, seq 4096, even routing), mean of iterations 11-20 of 20-iteration
runs; each row adds one change on top of the row above:
                                                    ms/iter TFLOP/s tok/s/GPU    mem  loss@20
  submitted FP8 tensorwise config (Turbo e2d9f1d7)   8894.0   670.6    29,474  277GB 11.33393
  + use_turbo_fused_act_with_probs (Primus flag)     7915.6   753.5    33,117  223GB 11.33435
  Turbo 1103b2df, base of these changes              7925.4   752.5    33,077  222GB 11.33427
  + flat tensorwise FP8 quant kernel                 7870.9   757.7    33,305  222GB 11.33415
  + per-token get_dispatch_layout                    7774.2   767.2    33,720  222GB 11.33431
  + HIP permute in DeepEPTokenDispatcher             7643.5   780.3    34,296  222GB 11.33423
  + turbo_deepep_num_cu 160 (Primus flag)            7353.4   811.1    35,649  223GB 11.33456
* + fused qk RMSNorm + RoPE, head_dim 128            6987.1   853.6    37,518  224GB 11.33373
This change (*): 7353.4 -> 6987.1 ms/iter (-4.98%),
35,649 -> 37,518 tokens/s/GPU (+5.24%).
Loss differences are within run-to-run noise (about 4e-4 at iteration 20).

Requires Primus-Turbo perf/flydsl/qk-rmsnorm-rope-hd128.
…DeepEP permute

New opt-in flag turbo_fp8_permute (default false). Under Primus-Turbo FP8
tensorwise current scaling with use_turbo_deepep and turbo_fused_grouped_gemm,
PrimusTurboDeepEPTokenDispatcher asks Turbo's DeepEPTokenDispatcher to
  - quantize the dispatched tokens before permute and return the permuted
    tokens as a grouped tensorwise QuantizedTensor
    (_post_dispatch(quantize_dtype=...)), and
  - return the experts' output gradient as a tensorwise QuantizedTensor from
    un-permute's backward (_pre_combine(grad_quantize_dtype=...)),
so the fused grouped MLP no longer quantizes the top-k times larger permuted
copies itself. The dtypes follow the FP8 format (HYBRID: e4m3 input, e5m2
gradient); results are bit-identical to the default path. Both dtypes are
derived in the call itself, so interleaved micro-batches need no dispatcher
state.

Any other recipe, or a layer with Turbo FP8 off, keeps the bf16 path.
Argument validation requires enable_primus_turbo, use_turbo_deepep and
turbo_fused_grouped_gemm, and rejects selective recompute of moe_act (that
path checkpoints the fused MLP and is not validated with QuantizedTensor
inputs and gradients).

Requires Primus-Turbo perf/moe/fp8-permute-tensorwise (moe_permute
quantize_dtype, moe_unpermute grad_quantize_dtype, QuantizedTensor grad_out in
grouped_mlp_fp8).

Qwen3-30B-A3B EP8 dispatcher call (32768 received tokens, hidden 2048,
16 local experts, top-8, pad_multiple 16), kernel time per call:
forward 645 -> 167 us, backward 687 -> 174 us.

End-to-end: Qwen3-30B-A3B FP8 tensorwise pretrain on 8x MI355X (EP8, MBS 8,
GBS 512, seq 4096, even routing), mean of iterations 11-20 of 20-iteration
runs; each row adds one change on top of the row above:
                                                    ms/iter TFLOP/s tok/s/GPU    mem  loss@20
  submitted FP8 tensorwise config (Turbo e2d9f1d7)   8894.0   670.6    29,474  277GB 11.33393
  + use_turbo_fused_act_with_probs (Primus flag)     7915.6   753.5    33,117  223GB 11.33435
  Turbo 1103b2df, base of these changes              7925.4   752.5    33,077  222GB 11.33427
  + flat tensorwise FP8 quant kernel                 7870.9   757.7    33,305  222GB 11.33415
  + per-token get_dispatch_layout                    7774.2   767.2    33,720  222GB 11.33431
  + HIP permute in DeepEPTokenDispatcher             7643.5   780.3    34,296  222GB 11.33423
  + turbo_deepep_num_cu 160 (Primus flag)            7353.4   811.1    35,649  223GB 11.33456
  + fused qk RMSNorm + RoPE, head_dim 128            6987.1   853.6    37,518  224GB 11.33373
  + turbo_fused_grouped_gemm (Primus flag)           6820.2   874.5    38,437  225GB 11.33412
* + turbo_fp8_permute                                6442.4   925.8    40,690  223GB 11.33404
Repeat runs of the last two rows: turbo_fused_grouped_gemm alone 6820.2 and
6815.1 ms; with turbo_fp8_permute 6442.4, 6515.0, 6523.8 and 6549.1 ms (the
later runs show more per-iteration jitter; the cause is not identified).
This change (*), means of those runs: 6817.7 -> 6507.6 ms/iter (-4.55%),
38,451 -> 40,283 tokens/s/GPU (+4.76%).
Loss differences are within run-to-run noise (about 4e-4 at iteration 20).

Tests: test_rocm_arg_validation.py covers validate_turbo_fp8_permute (each
required flag, selective recompute of moe_act, disabled no-op);
test_turbo_fp8_permute_dtypes.py checks _fp8_permute_dtypes returns
(e4m3, e5m2) for HYBRID and (e4m3, e4m3) for E4M3 tensorwise, and
(None, None) with the flag off, Turbo FP8 off or blockwise scaling.
…X config

Turn on every switch behind the 40k tokens/s/GPU Qwen3-30B-A3B FP8
tensorwise pretrain on 8x MI355X:
  - env PRIMUS_FUSED_QK_RMSNORM_ROPE=1 (top-level env:): fused packed-QKV
    split + Q/K RMSNorm + RoPE for head_dim 128;
  - turbo_fused_grouped_gemm and use_turbo_fused_act_with_probs;
  - turbo_fp8_permute: the DeepEP permute hands the fused grouped MLP FP8
    tokens and takes its FP8 output gradient;
  - turbo_deepep_num_cu 80 -> 160 (192 gives a NaN loss on one rank at
    iteration 5, 256 is slower).

turbo_fp8_permute needs a Primus-Turbo whose moe_permute takes
quantize_dtype (AMD-AGI/Primus-Turbo#555); the head_dim 128 fusion needs
AMD-AGI/Primus-Turbo#556 and falls back to the unfused path without it.

End-to-end: Qwen3-30B-A3B FP8 tensorwise pretrain on 8x MI355X (EP8, MBS 8,
GBS 512, seq 4096, even routing), mean of iterations 11-20 of 20-iteration
runs; each row adds one change on top of the row above:
                                                    ms/iter TFLOP/s tok/s/GPU    mem  loss@20
  FP8 tensorwise config (Turbo e2d9f1d7)             8894.0   670.6    29,474  277GB 11.33393
  + use_turbo_fused_act_with_probs                   7915.6   753.5    33,117  223GB 11.33435
  Turbo 1103b2df                                     7925.4   752.5    33,077  222GB 11.33427
  + flat tensorwise FP8 quant kernel (Turbo)         7870.9   757.7    33,305  222GB 11.33415
  + per-token get_dispatch_layout (Turbo)            7774.2   767.2    33,720  222GB 11.33431
  + HIP permute in DeepEPTokenDispatcher (Turbo)     7643.5   780.3    34,296  222GB 11.33423
  + turbo_deepep_num_cu 160                          7353.4   811.1    35,649  223GB 11.33456
  + fused qk RMSNorm + RoPE, head_dim 128            6987.1   853.6    37,518  224GB 11.33373
  + turbo_fused_grouped_gemm (mean of 2 runs)        6817.7   874.9    38,451  225GB 11.33412
  + turbo_fp8_permute (mean of 4 runs)               6507.6   916.7    40,283  223GB 11.33404
This branch on main 0a68e5c with the switches on: 6518.7 ms/iter with the
switches on the command line, 6410.6 and 6459.6 ms/iter from this config
alone (40,214 / 40,892 / 40,582 tokens/s/GPU).
Total: 8894.0 -> 6507.6 ms/iter (-26.8%), 29,474 -> 40,283 tokens/s/GPU
(+36.7%). Loss differences are within run-to-run noise (about 4e-4 at
iteration 20).

Known issue: one more config-only run hit a NaN forward loss on one rank
at iteration 3; its parsed arguments match a passing run exactly. Across
this work, 1 of 13 runs at 160 DeepEP CUs and the only run at 192 hit
this; none of the 5 runs at 80 CUs did, but those predate the fused paths,
so the cause is not isolated yet.
Copilot AI balanced review requested due to automatic review settings October 10, 2026 15:42

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

🟡 Changes recommended

The default recipe depends on unavailable Turbo APIs, permits an incompatible configuration, and enables a setting with known NaN failures.

3 open findings
What changed in this PR

Optimizes Qwen3-30B-A3B FP8 training using Primus-Turbo tensorwise FP8 and fused kernels.

Changes:

  • Adds FP8 DeepEP permute support and validation.
  • Expands fused Q/K RMSNorm + RoPE support to head dimension 128.
  • Enables and tunes Turbo optimizations in the MI355X recipe.
File Description
examples/​megatron/​configs/​MI355X/​qwen3_30B_A3B-FP8-pretrain.yaml Enables optimized Turbo FP8 recipe.
primus/​backends/​megatron/​core/​extensions/​primus_turbo.py Integrates FP8 attention and permute behavior.
primus/​backends/​megatron/​patches/​args/​rocm_arg_validation.py Validates FP8 permute configuration.
primus/​backends/​megatron/​patches/​turbo/​qk_rmsnorm_rope_patches.py Supports additional fused-attention shapes and norms.
primus/​configs/​modules/​megatron/​primus_turbo.yaml Defines the new opt-in flag.
tests/​unit_tests/​backends/​megatron/​test_rocm_arg_validation.py Tests argument validation.
tests/​unit_tests/​backends/​megatron/​test_turbo_fp8_permute_dtypes.py Tests FP8 dtype selection.

🧠 Review effort: Balanced


💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.

Comment on lines +8 to +9
env:
PRIMUS_FUSED_QK_RMSNORM_ROPE: "1"
Comment on lines +1140 to +1141
# flash_attn_fp8_func takes no ``sink`` argument.
sink_kwargs = {"sink": sink_tensor} if sink_tensor is not None else {}
option = "turbo_fp8_permute"
if not getattr(args, option, False):
return
for required in ("enable_primus_turbo", "use_turbo_deepep", "turbo_fused_grouped_gemm"):
Intermittent NaN forward loss at turbo_deepep_num_cu 160/192 (~1/13). The
cheap fence was only row-validated at <=80 CUs; DISABLE_CHEAP_FENCE selects
the release/acquire path until Turbo defaults high CU to full fence.
Copilot AI balanced review requested due to automatic review settings October 11, 2026 08:24

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

🟡 Changes recommended

The shipped Turbo pin lacks the enabled permutation API, and FP8 attention still fails when combined with sink attention.

5 open findings
Previously missed (1)

In code that hasn't changed since last review

Medium severity Test fused RoPE patching for head_dim 128 and TE RMSNorm

primus/​backends/​megatron/​patches/​turbo/​qk_rmsnorm_rope_patches.py:139

The existing test_qk_rmsnorm_rope_patches.py only tests DDP hook helpers; it does not exercise the newly broadened eligibility for head dimension 128 or plain TE RMSNorm. A regression that always falls back to the unfused path would therefore pass. Add a patch-level test with supported head_dim 128 and TE RMSNorm instances that verifies the fused operator is selected.

🧠 Review effort: Balanced

turbo_fused_grouped_gemm: true
use_turbo_fused_act_with_probs: true
# needs a Primus-Turbo whose moe_permute takes quantize_dtype
turbo_fp8_permute: true
# at 160/192 (~1/13). DISABLE=1 selects the release/acquire fence (~2-3% step).
env:
PRIMUS_FUSED_QK_RMSNORM_ROPE: "1"
PRIMUS_TURBO_DEEPEP_DISABLE_CHEAP_FENCE: "1"

This branch has not been deployed

No deployments
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