Skip to content

Commit a81c45e

Browse files
authored
Fix SM90 indexed A16 large-M scheduling (#65)
Signed-off-by: mgoin <mgoin64@gmail.com>
1 parent 8611853 commit a81c45e

2 files changed

Lines changed: 34 additions & 1 deletion

File tree

humming/tune/sm90_policies.py

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -115,7 +115,12 @@ def build_sm90_seed_config(problem: TuningProblem) -> dict:
115115
)
116116
if use_wide_indexed_tile:
117117
warp_shape_n = 64
118-
if layer_config.shape_k <= 512 and layer_config.shape_n >= 2048:
118+
# N=512 spills its accumulator at two-CTA residency from M=48 onward.
119+
if (
120+
layer_config.shape_k <= 512
121+
and layer_config.shape_n >= 2048
122+
and block_shape_m < 48
123+
):
119124
block_shape_n = 512
120125
block_shape_k = 64
121126
else:

tests/test_sm90_heuristics.py

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -288,6 +288,34 @@ def test_mxfp4_a16_indexed_preserves_warp_k_and_fits_grid(
288288
assert config["use_stream_k"] is expected_stream_k
289289

290290

291+
@pytest.mark.parametrize("shape_k", [512, 256])
292+
def test_mxfp4_a16_short_k_limits_wide_n_tile_before_block_m48(shape_k):
293+
layer = _layer(
294+
6144,
295+
shape_k,
296+
num_experts=256,
297+
a_dtype=dtypes.bfloat16,
298+
as_dtype=None,
299+
input_scale_group_size=0,
300+
)
301+
302+
small = Sm90Heuristics.get_config(
303+
layer,
304+
shape_m=6144,
305+
gemm_type=GemmType.INDEXED,
306+
)
307+
large = Sm90Heuristics.get_config(
308+
layer,
309+
shape_m=8192,
310+
gemm_type=GemmType.INDEXED,
311+
)
312+
313+
assert small["block_shape"] == (32, 512, 64)
314+
assert small["num_ctas_per_sm"] == 2
315+
assert large["block_shape"] == (48, 256, 64)
316+
assert large["num_ctas_per_sm"] == 2
317+
318+
291319
def test_nvfp4_a16_uses_narrow_first_wave_only():
292320
layer = _layer(
293321
5376,

0 commit comments

Comments
 (0)