Repository navigation
[Triton/Gluon] [Bugfix] Stop the split-K tail at K in the blockscale GEMMs - #5873
siliangchen-amd wants to merge 3 commits into
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
Critical kernel masking issues remain in both implementations, and the tests need to use supported configuration paths.
Review effort: Lite
Findings: 2
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_Khandling. - 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.
omuhamma
left a comment
There was a problem hiding this comment.
Have you considered adding a guard in the kernel loop instead? A guard that will prevent an 'empty' tile from corrupting the good values ?
|
@omuhamma Thanks, a guard in the loop is the better fix. Updated in 670827d:
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. |
There was a problem hiding this comment.
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
Open (5)
Resolved since last review (2)
|
@siliangchen-amd can you rebase? |
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.
19ff261 to
aed6c11
Compare
|
@azaidy Done. Thanks! |


Summary
compute_splitk_paramsroundsSPLITK_BLOCK_SIZEup to a multiple ofBLOCK_SIZE_Kand recomputesNUM_KSPLIT, soNUM_KSPLIT * SPLITK_BLOCK_SIZEcan exceed K: the last partition starts inside K and runs past it. The blockscale kernels still runSPLITK_BLOCK_SIZE / BLOCK_SIZE_KK blocks in every partition, andEVEN_Konly checksK % 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_preshuffledoes not round at all. A split that is not a multiple ofBLOCK_SIZE_Kstarts 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-MXFP8with SGLang at TP8 on MI325X, on aiter 0.1.19. The shared-expertdown_projhas 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 hitMemory 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
EVEN_SPLITKconstexpr (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_Kis unchanged.gemm_a8w8_blockscale_preshufflerounds its partitions withcompute_splitk_params, like the other three wrappers.compute_splitk_paramsbefore sizingyandy_pp, since it can lowerNUM_KSPLIT. The plain wrappers sized them first, soskip_reducereturned partial planes the kernel never wrote, and a config that dropped to one split launched withy=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_gemmat (16, 4096, 11008), preshuffle: gfx942's tuned config splits K into four 2816-wide partitions, 256 past Ktest_gemmat (M, 6144, 640 / 896), M in {1, 16, 64}, a16w8: gfx950's default config splits these into 256-wide partitions past Ktest_gemm_skip_reduce, a16w8 (16, 6144, 640): gfx950's default asks for 8 splits and normalizes to 3test_gemm_a8w8_blockscale.py+test_gemm_a16w8_blockscale.py, all casesAd hoc, with explicit configs that are not in the tests:
skip_reduce, 8 splits requested, K=640 (normalized to 3): all four wrappersskip_reduce, 8 splits requested, K=128 (normalized to 1): all four wrappersy=Noneat launchConfigs 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
_triton_kernelsimports.test_gemmadds (16, 4096, 11008), which aborts on main with gfx942's tuned preshuffle config.Found by Hyperloom; reviewed and measured by hand.
Hyperloom was optimizing
MiniMaxAI/MiniMax-M3-MXFP8on MI325X (TP8).