You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
{{ message }}
Repository navigation
perf(qwen3): Qwen3-30B-A3B FP8 on MI355X to 40k tokens/s/GPU with Primus-Turbo tensorwise - #1247
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:
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
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
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.
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.
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.
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.
# 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
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
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.
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:
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_permuteneeds a Primus-Turbo whosemoe_permutetakesquantize_dtype(AMD-AGI/Primus-Turbo#555). With an older Turbo, the first dispatch fails.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:
use_turbo_fused_act_with_probsget_dispatch_layoutDeepEPTokenDispatcherturbo_deepep_num_cu80 → 160turbo_fused_grouped_gemm(mean of 2 runs)turbo_fp8_permute(mean of 4 runs)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:
Known issue: intermittent NaN.
turbo_deepep_num_cu 160hit this, as did the only run at 192 (rank 2, iteration 5).Changes
fix(turbo): make enable_turbo_attention_float8 runnable againprimus/backends/megatron/core/extensions/primus_turbo.py:sink=only when a sink tensor exists;flash_attn_fp8_funchas nosinkparameter.flash_attn_fp8_funcpermutes sbhd storage as if it were the logical layout (fixed in Turbo byfix/fp8-attn-strided-layout).perf(qwen3): tune Qwen3-30B-A3B FP8 MI355X config to Turbo tensorwiseexamples/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 128primus/backends/megatron/patches/turbo/qk_rmsnorm_rope_patches.py:QK_RMSNORM_ROPE_HEAD_DIMS, falling back to(64,)on older Turbo.PrimusTurboRMSNorm. The fused path only reads weight / eps;zero_centered_gammais still rejected.perf(megatron): hand the Turbo fused grouped MLP FP8 tokens from the DeepEP permuteturbo_fp8_permute(default false inprimus_turbo.yaml).PrimusTurboDeepEPTokenDispatcherpassesquantize_dtypeto_post_dispatchandgrad_quantize_dtypeto_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.enable_primus_turbo,use_turbo_deepepandturbo_fused_grouped_gemm, and rejects selective recompute ofmoe_act.perf(qwen3): turn on the Turbo fusions in the Qwen3-30B-A3B FP8 MI355X configenv: PRIMUS_FUSED_QK_RMSNORM_ROPE: "1".turbo_fused_grouped_gemm,use_turbo_fused_act_with_probsandturbo_fp8_permuteset to true.turbo_deepep_num_cu80 → 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
tests/unit_tests/backends/megatron/test_rocm_arg_validation.py:validate_turbo_fp8_permute(each required flag, selective recompute ofmoe_act, disabled no-op).tests/unit_tests/backends/megatron/test_turbo_fp8_permute_dtypes.py:_fp8_permute_dtypesreturns (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.