def _select_cfg(M, N, K):
"""(BLOCK_M, BLOCK_N, BLOCK_K, num_warps, num_stages) — graph-tuned on gfx950.
The M-bucketed, shape-adaptive tile selection here is the speedup over the
upstream 2-bucket launcher. Tiles are pipelined (num_stages>=2, larger BLOCK_K)
and occupancy- and shape-aware: keyed on the LOCAL (M, N, K), so it adapts to the
TP-sharded shapes (e.g. MiniMax-M3 TP=4 vs TP=8, where local N and K differ) —
large-K prefill uses BLOCK_K=256; short-K (K=768) widens N. BLOCK_K must divide K
(the K-loop is unmasked), so every BLOCK_K below is guarded to be K-divisible
(served K: 384/768/1024/2048/6144).
"""
if M <= 64:
# decode (M in {1,32,64}): tiny-M GEMV is weight-BW + GPU-OCCUPANCY bound. The
# lever is NARROW BLOCK_N=16 (maximize N-tiles so more CUs stream the weight in
# parallel) + LARGE BLOCK_K (fewer K-iters, bigger coalesced weight loads).
# Tuned by CUDA-graph replay latency. Optimal at both TP=4 and TP=8.
if K % 1024 == 0: # K=2048, 6144 -> graph-best 16x16x1024 (all M)
return 16, 16, 1024, 2, 2
if K % 512 == 0:
return 16, 16, 512, 2, 3
if K % 256 == 0: # K=768 (shared_down) -> graph-best 16x32x256
return 16, 32, 256, 4, 3
return 16, 32, 128, 4, 3
# mid-M (65..256) on SMALL local-N: still occupancy-bound (a 64x64 tile makes too
# few N-tiles), so the narrow-BLOCK_N decode-style tile fills the CUs better.
# N<=1536 covers the real fused-qkv local N at TP=8: q heads shard but the GQA KV
# (4) + sparse-indexer (4) heads are < TP=8, so vLLM replicates them to 1/rank ->
# N = 1024 + 4*128 = 1536 (not 2560/2=1280). For the wider 1280<N<=1536 band the
# narrow tile only wins up to M=128 (at M=256 the 64x64 tile is better), so cap it
# there; N<=1280 keeps the narrow tile through M=256. TP=4 qkv N=2560 is unchanged.
if (M <= 256 and N <= 1280) or (M <= 128 and N <= 1536):
if K % 1024 == 0:
return 16, 16, 1024, 2, 2
if K % 512 == 0:
return 16, 16, 512, 2, 3
if K % 256 == 0:
return 16, 16, 256, 2, 3
return 16, 16, 128, 2, 3
# right-sized launch grid (host-side), used to gate the tall 256-BLOCK_M tile.
occ = triton.cdiv(M, 256) * triton.cdiv(N, 128)
if K <= 1024: # short-K (shared_down K=384/768; TP=8 o_proj K=1024)
if M <= 256:
return (64, 64, 256, 8, 2) if K % 256 == 0 else (64, 64, 128, 8, 2)
# large prefill. (The former 256x128x256 tall tile was faster only on triton
# 3.6; on triton 3.7 its large BLOCK_M register/LDS footprint spills or hits
# "out of resources", so use 128x128x256 -- within the known-good footprint.)
if M >= 4096 and K >= 1024 and K % 256 == 0 and occ >= 256:
return 128, 128, 256, 8, 3
return 128, 128, 128, 8, 3
# large-K (K >= 2048). BLOCK_K is K-divisibility-guarded (the K-loop is unmasked):
# served large-K is 2048/6144 (%256==0), but fall back to 128 (always divides, since
# the entry requires K%128==0) for any other K to stay correct.
if M <= 256: # conc~128 decode + small prefill chunk: occupancy tile
if K % 512 == 0:
return 64, 64, 512, 8, 2
return (64, 64, 256, 8, 2) if K % 256 == 0 else (64, 64, 128, 8, 2)
if M <= 1024: # medium prefill chunk: BN=64 keeps small-N occupied
return (128, 64, 256, 8, 3) if K % 256 == 0 else (128, 64, 128, 8, 3)
# large prefill (M > 1024). The previously graph-tuned tall 256x128x256 and deep
# 128x128x512 tiles won only on triton 3.6; on triton 3.7 their larger BLOCK_M /
# BLOCK_K register+LDS footprint spills (or hits "out of resources" on stricter
# ROCm/triton builds). The 128x128x256 tile is equal-or-faster on triton 3.7 (the
# M=4096,N=2048,K=6144 shape: ~104us vs the 256x128x256 tile's ~184us), within ~5%
# on 3.6, and inside the footprint of the tiles used elsewhere in this selector.
# Covers the qkv-class local N=1536 (TP=8 qkv / TP=4 shared_gate_up) and the deep-K
# / very-large-M shapes.
if K % 256 == 0 and (1280 < N <= 1536 or (occ >= 128 and (K >= 4096 or M >= 4096))):
return 128, 128, 256, 8, 3
# small local-N (e.g. TP=8 shared_gate_up N=768): a 64-wide BLOCK_N doubles the
# N-tile count -> better CU fill than 128x128 at this mid-large M (~1.4x there).
if N <= 1024 and K % 256 == 0:
return 128, 64, 256, 8, 3
return (128, 128, 256, 8, 2) if K % 256 == 0 else (128, 256, 128, 8, 3)