Repository navigation
[CK] [Bugfix] Keep two K tiles per split in blockscale MoE stage-1 split-K - #5979
Conversation
🏷️ CI GuideRuns automatically on every PR:
Extended tests (opt-in via labels):
One backend per PR: PR title tags & labels: |
There was a problem hiding this comment.
Copilot review overview
🟡 Changes recommended
The shared gfx950 path is excluded from the new architecture-sensitive regression test.
Review effort: Balanced
Findings: 1
Open (1)
What changed in this PR
Updates CK blockscale MoE stage-1 split-K handling to prevent unsafe single-tile splits.
Changes:
- Recalculates
KBatchto retain at least two K tiles per split. - Adds correctness and error-path regression tests.
| File | Description |
|---|---|
gemm_moe_ck2stages_common_blockscale.cuh |
Safely adjusts undersized split-K partitions. |
test_moe_blockscale_ksplit.py |
Tests split adjustment and single-tile rejection. |
💡 Configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
…it-K The default fp8-blockscale 2-stage fused MoE returns wrong output whenever get_ksplit() asks for model_dim / 256 splits, one 256-wide K tile each. This happens at small decode batches on gfx942 (cu_num 304), for example at one token for (model_dim, inter_dim) = (4096, 256) with 288 experts and top-8, (7168, 128) with 257 experts and top-9, and (2048, 384) with 128 experts and top-8. Every output element differs from the reference, and the values change from run to run. ck_moe_stage1_gemm turns the requested number of splits into CK's KBatch, the number of KPerBlock tiles each split covers: K / (splitk * KPerBlock). With one tile per split that is 1, which DeviceMoeGemmBlockScale reads as "no split": it runs a single z-block over the whole K and skips zeroing the output. The split-K instance still accumulates into that output atomically, and fused_moe allocates it with torch.empty. StrideE also keys the split-K layout (2 * N per row) on KBatch > 1, so the rows overlap as well. ROCm#3997 fixed this in ck_moe_stage1() by turning split-K off for such requests, and the MoE refactor in ROCm#4394 dropped that guard. This fix sits at the KBatch conversion instead, so no caller reaches KBatch == 1 through split-K, whether the split comes from get_ksplit(), a tuned CSV or AITER_KSPLIT. It also keeps the split rather than turning it off: when splitk would leave fewer than two tiles per split, stage 1 uses fewer splits of at least two tiles, KBatch = 2 or the smallest divisor of the tile count that is at least 2. At the (4096, 256, E=288, top-8) one-token shape the operator then takes 32.1 us, against 41.1 us without split-K. A K of a single tile cannot be split and now fails with an error. Requests that already give two or more tiles per split are unchanged. Signed-off-by: Jin Tao <jintao12@amd.com>
9a70f46 to
75c2199
Compare
|
I tested this on my end and this fixes gibberish output from glm-5.3-flash at concurrency 1 on gfx950 in vllm |

Motivation
The default fp8-blockscale 2-stage
fused_moereturns garbage at small decode batches wheneverget_ksplit()picksmodel_dim / 256splits, one 256-wide K tile per split. On gfx942 (MI325X,cu_num304) this includes, fromop_tests/test_moe_2stage.py -q 5:get_ksplitinfNaNThe values change from run to run. This is a regression: #3997 fixed it on July 15, and the MoE refactor in #4394 dropped that guard on July 31.
Technical Details
ck_moe_stage1_gemmturns the requested number of splits into CK'sKBatch, the number ofKPerBlocktiles each split covers:K / (splitk * KPerBlock). The default 16-row stage-1 instance hasKPerBlock256, so one tile per split givesKBatch == 1, andDeviceMoeGemmBlockScalereads that as "no split":hipMemsetAsyncof the output, which it only issues whenKBatch > 1;ck_moe_stage1()allocates it withtorch.empty;StrideEin the wrapper also keys the split-K layout (2 * N per row) onKBatch > 1, so rows overlap as well.This PR makes the wrapper keep at least two tiles per split. When
splitkwould leave fewer, stage 1 uses fewer splits instead:KBatch = 2, or the smallest divisor of the tile count that is at least 2 (nine tiles become three splits of three, seven become one split of seven). A K of a single tile cannot be split, and now fails with an error instead of returning garbage. Requests that already give two or more tiles per split are unchanged.Why here rather than restoring #3997's guard in
ck_moe_stage1():KBatch == 1, so no caller can reach it, whether the split comes fromget_ksplit(), a tuned CSV orAITER_KSPLIT, and a Python refactor cannot drop it again.Test Plan
op_tests/test_moe_2stage.py -q 5over 19 shapes, before and after:-dim 4096,256 -e 288 -k 8 -t 1 2 4 16 32,-dim 7168,128 7168,256 -e 257 -k 9 -t 1 3 5 16,-dim 2048,384 2048,768 -e 128 -k 8 -t 1 2 4.AITER_KSPLITsweep, before and after, comparingfused_moewith an fp32 reference built from the dequantized weights: one tile per split at model_dim 4096 (routed and single-expert), 2048 and 8192, nine and seven tiles, and a single-tile K that must raise.gemm_moe_tune.py --run_config, two runs each.Test Result
test_moe_2stage.py -q 5: the three one-tile cases in the table above go from every element wrong to the error level of the other token counts of the same shape. Their max error becomes 1280 (1024-1408 at the other token counts), 1024 (1536) and 624 (768). The other 16 cases give the same errors as before.Forced split, cosine to the fp32 reference at one token:
The gap to 1.0 is activation quantization, which the reference omits. A single-tile K (model_dim 256, forced split 2) now raises
K(256) must span at least two KPerBlock(256) tiles for split-K.Production operator at one token, (4096, 256), mean of two runs:
Shapes where the split already had two or more tiles time the same as before, for example 39.4-40.1 µs at two tokens.
#5500 and #5953, which add GLM-5.3-Flash fused-MoE tables for gfx942, leave this TP8 one-token shape to the default and rely on this fix; the best tuned candidate there takes 59.8 µs.
Environment
AMD Instinct MI325X (
gfx942,cu_num304), containeramdsiloai/vllm:vllm-openai-rocm-glm5.3-flash-mi325-28092026-pr57161(ROCm 7.2.3, HIP 7.2.53211, PyTorch 2.12.0), one GPU. Tested on aiter main2e6209429with the split-K module rebuilt; the changed file andaiter/fused_moe.pyare identical at89a47b84a, which this PR is based on.Submission Checklist
main).