Skip to content

[Triton/Gluon] [Bugfix] Stop the split-K tail at K in the blockscale GEMMs - #5873

Open
siliangchen-amd wants to merge 3 commits into
ROCm:mainfrom
siliangchen-amd:fix-blockscale-splitk-even-k
Open

siliangchen-amd wants to merge 3 commits into
ROCm:mainfrom
siliangchen-amd:fix-blockscale-splitk-even-k

Conversation

@siliangchen-amd

@siliangchen-amd siliangchen-amd commented Sep 26, 2026 •

Copy link
Copy Markdown
Contributor

Summary

compute_splitk_params rounds SPLITK_BLOCK_SIZE up to a multiple of BLOCK_SIZE_K and recomputes NUM_KSPLIT, so NUM_KSPLIT * SPLITK_BLOCK_SIZE can exceed K: the last partition starts inside K and runs past it. The blockscale kernels still run SPLITK_BLOCK_SIZE / BLOCK_SIZE_K K blocks in every partition, and EVEN_K only checks K % BLOCK_SIZE_K == 0, so the tile past K is loaded without masks, along with its block scales. The output is wrong or NaN, and the stray reads can fault.

gemm_a8w8_blockscale_preshuffle does not round at all. A split that is not a multiple of BLOCK_SIZE_K starts in the middle of a block, so the overlap is added twice and the last split still runs past K. gfx942's tuned preshuffle table has such a config: N=4096, K=11008, NUM_KSPLIT=4 (2752-wide splits) for M <= 128, and it aborts with a GPU memory fault on main.

We hit this serving MiniMaxAI/MiniMax-M3-MXFP8 with SGLang at TP8 on MI325X, on aiter 0.1.19. The shared-expert down_proj has K=384 per rank, and 0.1.19's default config splits it into two 256-wide partitions, so the model decoded token 0 and CUDA-graph capture hit Memory access fault by GPU. Main's default gfx942 config no longer splits that shape, but any config whose partitions overshoot K still does this.

Fix

  • When the partitions overshoot K, the K loop ends at the last K block, so the empty tile never loads A, B or its block scales. A new EVEN_SPLITK constexpr (NUM_KSPLIT * SPLITK_BLOCK_SIZE == K) selects the bound, so configs whose partitions end at K compile exactly as before. This covers all four kernels: a8w8 and a16w8, plain and preshuffle. EVEN_K is unchanged.
  • gemm_a8w8_blockscale_preshuffle rounds its partitions with compute_splitk_params, like the other three wrappers.
  • All four wrappers call compute_splitk_params before sizing y and y_pp, since it can lower NUM_KSPLIT. The plain wrappers sized them first, so skip_reduce returned partial planes the kernel never wrote, and a config that dropped to one split launched with y=None.

Results (MI325X, gfx942, triton backend)

The tests use the public wrappers and their JSON configs only, so they reach the tail through shipped configs:

test main this PR
test_gemm at (16, 4096, 11008), preshuffle: gfx942's tuned config splits K into four 2816-wide partitions, 256 past K GPU memory fault, process aborts pass
test_gemm at (M, 6144, 640 / 896), M in {1, 16, 64}, a16w8: gfx950's default config splits these into 256-wide partitions past K pass (gfx942's default does not split them; gfx950 CI exercises the tail)
test_gemm_skip_reduce, a16w8 (16, 6144, 640): gfx950's default asks for 8 splits and normalizes to 3 pass
test_gemm_a8w8_blockscale.py + test_gemm_a16w8_blockscale.py, all cases 279 passed, 250 skipped

Ad hoc, with explicit configs that are not in the tests:

check before after
forced 8-way split, K in {640, 896}: all four kernels 8 fail on main 8 pass
skip_reduce, 8 splits requested, K=640 (normalized to 3): all four wrappers 8 planes returned, 3 written 3 planes, correct
skip_reduce, 8 splits requested, K=128 (normalized to 1): all four wrappers y=None at launch runs, correct

Configs whose partitions end at K keep their timing: 23 default-config shapes (a8w8 and a16w8, plain and preshuffle, M from 1 to 1024) are within 1.2% of main.

Test plan

  • Tests go through the public wrappers and shipped JSON configs, with no config overrides or _triton_kernels imports.
  • test_gemm adds (16, 4096, 11008), which aborts on main with gfx942's tuned preshuffle config.
  • Full a8w8 / a16w8 blockscale test files: 279 passed, 250 skipped (gluon on gfx942, small-K triton shapes).
  • black and ruff 0.16.0.
  • Only validated on gfx942 with the triton backend; the a16w8 tail shapes split only with gfx950's default config.

Found by Hyperloom; reviewed and measured by hand.
Hyperloom was optimizing MiniMaxAI/MiniMax-M3-MXFP8 on MI325X (TP8).

@siliangchen-amd
siliangchen-amd requested a review from a team September 26, 2026 14:20
@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; main branch always runs both MI35X and MI300X
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 5873 --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.

@github-actions github-actions Bot changed the title [Triton] [Bugfix] Mask the split-K tail in the blockscale GEMMs [Triton/Gluon] [Bugfix] Mask the split-K tail in the blockscale GEMMs Sep 26, 2026
@Boss2002n
Boss2002n requested a lite review from Copilot September 26, 2026 17:33

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

Critical kernel masking issues remain in both implementations, and the tests need to use supported configuration paths.

Review effort: Lite
Findings: 2 High severity · 2 Low severity

Open (4)
What changed in this PR

Fixes split-K tail handling in Triton blockscale GEMMs by updating masking heuristics and adding regression tests.

Changes:

  • Updates A8W8 and A16W8 EVEN_K handling.
  • Adds forced split-K and non-aligned-K test coverage.
File Summary
op_tests/​triton_tests/​gemm/​basic/​test_gemm_a8w8_blockscale.py Adds split-K tail tests; uses private configuration helpers and literal tuning overrides.
op_tests/​triton_tests/​gemm/​basic/​test_gemm_a16w8_blockscale.py Adds A16W8 split-K coverage with the same configuration concerns.
aiter/​ops/​triton/​_triton_kernels/​gemm/​basic/​gemm_a8w8_blockscale.py Updates tail masking, but scale loads remain unmasked for padded iterations.
aiter/​ops/​triton/​_triton_kernels/​gemm/​basic/​gemm_a16w8_blockscale.py Enables the predicate without providing the required masked-load path.

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

Comment thread aiter/ops/triton/_triton_kernels/gemm/basic/gemm_a16w8_blockscale.py Outdated
Comment thread aiter/ops/triton/_triton_kernels/gemm/basic/gemm_a8w8_blockscale.py Outdated
Comment thread op_tests/triton_tests/gemm/basic/test_gemm_a16w8_blockscale.py Outdated
Comment thread op_tests/triton_tests/gemm/basic/test_gemm_a8w8_blockscale.py Outdated
@zufayu
zufayu requested review from a team and vgokhale September 27, 2026 00:33
@azaidy
azaidy requested a review from omuhamma September 28, 2026 15:27
@azaidy azaidy assigned omuhamma and unassigned vgokhale Sep 28, 2026

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

Have you considered adding a guard in the kernel loop instead? A guard that will prevent an 'empty' tile from corrupting the good values ?

@siliangchen-amd siliangchen-amd changed the title [Triton/Gluon] [Bugfix] Mask the split-K tail in the blockscale GEMMs [Triton/Gluon] [Bugfix] Stop the split-K tail at K in the blockscale GEMMs Sep 29, 2026
@siliangchen-amd

Copy link
Copy Markdown
Contributor Author

@omuhamma Thanks, a guard in the loop is the better fix. Updated in 670827d:

  • When the partitions overshoot K, the K loop now ends at the last K block, so the empty tile never loads A, B or its block scales (masking A and B alone didn't cover the scales). A new EVEN_SPLITK constexpr selects the bound, so configs whose partitions end at K compile exactly as before: 23 default-config shapes are within 1.2% of main on MI325X.
  • The bound doesn't need a masked path, so the two preshuffle kernels get it too. While testing them I found that gemm_a8w8_blockscale_preshuffle doesn't round its partitions at all, so it now calls compute_splitk_params like the other wrappers. Without that, gfx942's tuned N=4096, K=11008 preshuffle config aborts with a GPU memory fault on main.
  • test_gemm_splitk_tail now covers all four kernels (fails on main, passes here), and the full a8w8 / a16w8 blockscale test files pass.

One trade-off: on forced overshooting configs, the bounded loop ranges from 1.5x slower than masking every tile (2 K blocks per split) to 1.65x faster (13 blocks per split), because the trip count is no longer a compile-time constant. Those configs return wrong results on main today, and the only shipped per-shape table that has one is the preshuffle entry above.

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

Split normalization occurs after output allocation, potentially returning uninitialized partials when the split count decreases.

Review effort: Balanced
Findings: 1 High severity · 4 Low severity

Open (5)
Resolved since last review (2)

Comment thread aiter/ops/triton/gemm/basic/gemm_a8w8_blockscale.py Outdated
Comment thread op_tests/triton_tests/gemm/basic/test_gemm_a16w8_blockscale.py Outdated
Comment thread op_tests/triton_tests/gemm/basic/test_gemm_a8w8_blockscale.py Outdated
@azaidy

azaidy commented Oct 7, 2026

Copy link
Copy Markdown
Contributor

@siliangchen-amd can you rebase?

siliangchen-amd and others added 3 commits October 9, 2026 08:31
compute_splitk_params rounds SPLITK_BLOCK_SIZE up to a multiple of
BLOCK_SIZE_K, so NUM_KSPLIT * SPLITK_BLOCK_SIZE can exceed K and the last
partition runs past the end of A and B. EVEN_K only checked
K % BLOCK_SIZE_K, so those tiles were loaded unmasked: the result is wrong
(or NaN) and the stray reads can fault.

On gfx942 the default A8W8 config splits K eight ways for M <= 64, which
hits this for K = 384, 640 or 896. Require K % SPLITK_BLOCK_SIZE == 0 as
well in the a8w8 and a16w8 kernels. The preshuffle kernels are left alone:
their wrappers size the split differently, and the a16w8 one has no masked
path at all.

Signed-off-by: siliangchen-amd <SiLiang.Chen@amd.com>
…y tile

SPLITK_BLOCK_SIZE is rounded up to BLOCK_SIZE_K, so the last partition can
run past K. Instead of turning EVEN_K off for every tile, the K loop now
ends at the last K block when the partitions overshoot K (EVEN_SPLITK is
false). The empty tile then never loads A, B or its block scales, which
masking A and B alone did not cover, and the preshuffle kernels, which
have no masked path for B, get the same bound. Configs whose partitions
end exactly at K compile as before.

gemm_a8w8_blockscale_preshuffle now rounds its partitions with
compute_splitk_params like the other three wrappers. Without it, a split
that is not a multiple of BLOCK_SIZE_K starts in the middle of a block
and adds the overlap twice; gfx942's tuned N=4096, K=11008 preshuffle
config (four 2752-wide splits) aborts with a GPU memory fault on main.

Co-authored-by: Cursor <cursoragent@cursor.com>
compute_splitk_params can lower NUM_KSPLIT, but the four blockscale
wrappers sized y and y_pp from the requested value first. With
skip_reduce the caller got partial planes the kernel never wrote, and
when the split count dropped to 1, y was still None at launch. Call it
before allocating.

Test the split-K tail through the wrappers' own config resolution
instead of a hardcoded NUM_KSPLIT: gfx942's tuned preshuffled a8w8
config overshoots K=11008, and gfx950's a16w8 default overshoots K=640
and 896. A skip_reduce test covers the lowered split count.
@siliangchen-amd
siliangchen-amd force-pushed the fix-blockscale-splitk-even-k branch from 19ff261 to aed6c11 Compare October 9, 2026 08:34
@siliangchen-amd

Copy link
Copy Markdown
Contributor Author

@azaidy Done. Thanks!

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

Projects

None yet

Development

Successfully merging this pull request may close these issues.

5 participants