Skip to content

[CK] [Bugfix] Keep two K tiles per split in blockscale MoE stage-1 split-K - #5979

Merged
yifehuan merged 1 commit into
ROCm:mainfrom
jin-amd:fix-ck-moe-blockscale-splitk-one-tile
Oct 2, 2026
Merged

yifehuan merged 1 commit into
ROCm:mainfrom
jin-amd:fix-ck-moe-blockscale-splitk-one-tile

Conversation

@jin-amd

@jin-amd jin-amd commented Sep 30, 2026 •

Copy link
Copy Markdown
Contributor

Motivation

The default fp8-blockscale 2-stage fused_moe returns garbage at small decode batches whenever get_ksplit() picks model_dim / 256 splits, one 256-wide K tile per split. On gfx942 (MI325X, cu_num 304) this includes, from op_tests/test_moe_2stage.py -q 5:

model_dim, inter_dim experts, topk tokens get_ksplit before this PR
4096, 256 (GLM-5.3-Flash at TP8) 288, 8 1 16 every element wrong, max error inf
7168, 128 257, 9 1 28 every element wrong, max error 1.8e6
2048, 384 128, 8 1 8 every element wrong, NaN

The 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_gemm turns the requested number of splits into CK's KBatch, the number of KPerBlock tiles each split covers: K / (splitk * KPerBlock). The default 16-row stage-1 instance has KPerBlock 256, so one tile per split gives KBatch == 1, and DeviceMoeGemmBlockScale reads that as "no split":

  • it launches a single z-block over the whole K and skips the hipMemsetAsync of the output, which it only issues when KBatch > 1;
  • the split-K instance still accumulates into that output atomically, and ck_moe_stage1() allocates it with torch.empty;
  • StrideE in the wrapper also keys the split-K layout (2 * N per row) on KBatch > 1, so rows overlap as well.

This PR makes the wrapper keep at least two tiles per split. When splitk would 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():

  • The guard sits next to the conversion that collides with CK's KBatch == 1, so no caller can reach it, whether the split comes from get_ksplit(), a tuned CSV or AITER_KSPLIT, and a Python refactor cannot drop it again.
  • It keeps the split. fix(fused_moe): require KBatch >= 2 for block-fp8 split-k #3997 turned split-K off for these requests. At the GLM-5.3-Flash TP8 one-token shape, the operator takes 32.1 µs with this PR against 41.1 µs with split-K off.

Test Plan

  • op_tests/test_moe_2stage.py -q 5 over 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.
  • A forced-AITER_KSPLIT sweep, before and after, comparing fused_moe with 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.
  • Production-operator timing with 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:

model_dim splits K per split before after
4096 8 512 0.99895 0.99895
4096 16 256 0.884 0.99895
4096 (single expert) 16 256 0.771 0.99908
2048 8 256 0.936 0.99929
8192 32 256 0.848 0.99909
2304 9 256 0.823 0.99910
1792 7 256 0.625 0.99889

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:

routed, 288 experts, top-8 fused shared, 289 experts, top-9
before this PR (wrong output) 36.7 µs 37.5 µs
split-K off, as in #3997 41.1 µs 42.5 µs
this PR 32.1 µs 34.9 µs

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_num 304), container amdsiloai/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 main 2e6209429 with the split-K module rebuilt; the changed file and aiter/fused_moe.py are identical at 89a47b84a, which this PR is based on.

Submission Checklist

  • Looked over the ROCm contributing guidelines.
  • Targeting the repository default branch (main).
  • Commit includes the DCO sign-off.
  • CI green.

@jin-amd
jin-amd requested review from a team and a balanced review from Copilot September 30, 2026 05:49
@github-actions github-actions Bot changed the title [Bugfix][CK] Keep two K tiles per split in blockscale MoE stage-1 split-K [CK] [Bugfix] Keep two K tiles per split in blockscale MoE stage-1 split-K Sep 30, 2026
@github-actions

Copy link
Copy Markdown
Contributor

🏷️ CI Guide

Runs automatically on every PR:

  • ✅ Pre-checks (submodule verification, code formatting)
  • ✅ Aiter op tests (gfx942 + gfx950)
  • ✅ Triton tests on MI35X (only when aiter/ops/triton/** or related paths are changed)

Extended tests (opt-in via labels):

Label Tests
ci:gfx1250-ffm-triton Run the five-shard gfx1250 FFM Triton test suite
ci:triton-300x Run an additional Triton test job on MI300X in PRs (added automatically when gfx942 configs change); main branch always runs both MI35X and MI300X
ci:triton-355 Run the full Triton test suite on MI35X, not only the tests the change affects
multigpu Aiter multi-GPU tests on the 8-GPU runner
ci:sglang SGLang integration tests: DeepSeek-R1-MXFP4 accuracy, Qwen 3.5 accuracy
ci:atom ATOM benchmark: DeepSeek-R1-0528, GPT-OSS-120B
ci:atom_full ATOM accuracy suite for PR and main models from ATOM models_accuracy.json
ci:vllm vLLM benchmark: GPT-OSS-120B, DeepSeek-R1-0528, Kimi-K2.5
ci:all All standard extended tests (excludes ci:atom_full)

Only add ci:atom_full for FlyDSL or Triton upgrades.
Add labels via the sidebar or gh pr edit 5979 --add-label <label>

One backend per PR:
A PR changes one kernel backend: [Triton/Gluon] (Triton and Gluon count as one), [HIP], [ASM], [CK], [OPUS] or [FlyDSL]. If the title ends up with two backend tags, split the PR -- as stacked pull requests when one part cannot merge without the other.

PR title tags & labels:
Component tags ([Triton/Gluon], [HIP], [CK], [ASM], ...) are added to the PR title and as PR labels automatically from the changed files and re-synced on every push — change-type tags like [fix]/[Perf], op tags like [MLA], and human labels (ci:*) are left untouched. Add the no-auto-title label to stop the title rewrites; labels stay in sync either way.

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.

Copilot review overview

🟡 Changes recommended

The shared gfx950 path is excluded from the new architecture-sensitive regression test.

Review effort: Balanced
Findings: 1 Medium severity

Open (1)
What changed in this PR

Updates CK blockscale MoE stage-1 split-K handling to prevent unsafe single-tile splits.

Changes:

  • Recalculates KBatch to 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.

Comment thread op_tests/test_moe_blockscale_ksplit.py Outdated
…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>
Copilot AI balanced review requested due to automatic review settings September 30, 2026 05:53
@jin-amd
jin-amd force-pushed the fix-ck-moe-blockscale-splitk-one-tile branch from 9a70f46 to 75c2199 Compare September 30, 2026 05:53

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.

Copilot review overview

🟡 Changes recommended

The regression test claimed in the PR description is missing from the changes.

Review effort: Balanced
Findings: 2 Medium severity

Open (2)

@simondanielsson

Copy link
Copy Markdown
Contributor

I tested this on my end and this fixes gibberish output from glm-5.3-flash at concurrency 1 on gfx950 in vllm

@yifehuan yifehuan 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.

LGTM

@yifehuan
yifehuan merged commit 635c10f into ROCm:main Oct 2, 2026
63 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants