Skip to content

fix(tune): restore late AS and split issue/wait for M-major W4A8 grouped GEMM - #126

Merged
jinzhen-lin merged 2 commits into
vllm-project:mainfrom
LopezCastroRoberto:fix/sm90-w4a8-grouped-late-as
Oct 9, 2026
Merged

jinzhen-lin merged 2 commits into
vllm-project:mainfrom
LopezCastroRoberto:fix/sm90-w4a8-grouped-late-as

Conversation

@LopezCastroRoberto

Copy link
Copy Markdown
Contributor

fix(tune): restore late AS and split issue/wait for M-major W4A8 grouped GEMM

Summary

#121 made late activation-scale promotion and split WGMMA issue/wait opt-in tuning flags (wgmma_use_late_as, wgmma_split_issue_wait, both default False). Before #121, both behaviors were implied at compile time by kUsePackedLateAS for exactly the M-major grouped-contiguous MXFP4 (GS32) × FP8 (GS128) schedule added in #95. _set_w4a8_config was not updated, so from #121 on that schedule compiles with both disabled, and grouped M-major prefill regressed by ~20% at M ≥ 2048.

This PR sets both flags in _set_w4a8_config:

  • wgmma_use_late_as=True for every tile.
  • wgmma_split_issue_wait=True for tiles up to M160. M176 exceeds the split issue/wait register budget (accumulator 176 + single-buffer 16 + input scales 44 = 236 >= 224); M160 fits. Late AS without split is slower than neither flag on M144/M160 tiles, so the split cutoff matters.

There are no CUDA changes. Output is bit-identical to the flags-off path on every tile size the policy selects (64, 144, 160, 176) for N4096/K6144 and N6144/K2048.

test_grouped_w4a8_ranges_match_direct_selection now asserts both flags on every selected range and that each selected config passes get_register_budget_error. The updated test fails on current main (16/16) and passes with this change.

Where the regression comes from

M-major grouped-contiguous, N4096/K6144, 32 experts, top-k 8, balanced, latency in ms:

Commit M=2048 M=8192
763c3b7 (parent of #121) 1.054 3.994
1f212e2 (#121) 1.274 4.872
1f212e2 + wgmma_use_late_as 1.066 4.034
1f212e2 + both flags 1.066 4.037

No other commit between #95 and main moves these numbers.

Benchmarks

H200, benchmarks/bench_humming.py, main (f40f6ec) vs this branch, latency in ms. 32 experts, top-k 8, default (unbalanced) routing, M-major input scales. Each value is the median of 3 runs, alternating between the two branches.

python benchmarks/bench_humming.py --shape_n 4096 --shape_k 6144 \
  --a_dtype float8e4m3 --b_dtype float4e2m1 --c_dtype bfloat16 --bs_dtype float8e8m0 \
  --input_scale_group_size 128 --weight_scale_group_size 32 \
  --gemm_type grouped_contiguous --num_experts 32 --top_k 8 --use_m_major_input_scale

Gate/up (N=4096, K=6144)

M main this PR speedup
16 0.2091 0.2100 1.00×
64 0.2177 0.2178 1.00×
256 0.2962 0.2609 1.14×
512 0.3251 0.3162 1.03×
1024 0.6115 0.5974 1.02×
2048 1.3364 1.1170 1.20×
4096 2.5421 2.1021 1.21×
16384 9.6990 7.9706 1.22×

Down (N=6144, K=2048)

M main this PR speedup
16 0.1137 0.1137 1.00×
64 0.1159 0.1159 1.00×
256 0.1360 0.1349 1.01×
512 0.2071 0.1779 1.16×
1024 0.3839 0.3387 1.13×
2048 0.7356 0.5940 1.24×
4096 1.3767 1.1695 1.18×
16384 5.2205 4.4217 1.18×

Only the M-major grouped-contiguous path changes; other GEMM types and row-major scales do not use _set_w4a8_config.

Split cutoff on M160 tiles (gate/up, balanced, ms):

M (tile) split ≤ 144 split ≤ 160 before #121
576 (M160) 0.3551 0.3394 0.3383
1152 (M160) 0.6831 0.6456 0.6412

Tests

All run on H200 against this branch:

  • pytest tests/test_sm90_heuristics.py tests/test_tuning_candidates.py: 77 passed, 99 skipped
  • pytest tests/kernels/humming/test_moe.py tests/kernels/humming/test_mxfp4.py tests/kernels/humming/test_special_paths.py: 1538 passed, 12 skipped
  • ruff check and ruff format --check clean on the changed files

Note on #124

#124 sets the same two flags in _set_w4a8_config as part of a broader change, with the split cutoff at M144. This PR is the minimal fix for the #121 regression. With an M144 cutoff, M160 tiles stay ~4-6% slower than before #121 (table above), so whichever PR lands second should use the M160 cutoff.

🤖 Generated with Claude Code

…ped GEMM

vllm-project#121 made late activation-scale promotion and split WGMMA issue/wait
opt-in tuning flags that default to off. Before vllm-project#121 both were implied
by kUsePackedLateAS for the M-major grouped-contiguous MXFP4 x FP8_BLOCK
schedule from vllm-project#95, but _set_w4a8_config was never updated, so that
schedule regressed by ~20% at M >= 2048.

Set wgmma_use_late_as on every tile and wgmma_split_issue_wait up to
M160. M176 exceeds the split register budget; late AS without split is
slower on M144/M160 tiles. Output is bit-identical to the flags-off path.

Extend test_grouped_w4a8_ranges_match_direct_selection to assert both
flags and that every selected config passes the register budget check.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Signed-off-by: LopezCastroRoberto <rocastro@redhat.com>

@jinzhen-lin jinzhen-lin left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thank you! I missed this before.

Comment thread humming/tune/sm90_policies.py Outdated
Signed-off-by: LopezCastroRoberto <rocastro@redhat.com>
@jinzhen-lin
jinzhen-lin merged commit d113632 into vllm-project:main Oct 9, 2026
4 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants