From fae0e22cc6635cd63540e9069652dbe4542ed530 Mon Sep 17 00:00:00 2001 From: Justin Hu Date: Fri, 14 Aug 2026 01:14:53 +0000 Subject: [PATCH 1/2] perf: tune SM100 GEMM scheduling for dW shapes Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- .../ops/cutedsl/ops/_sm100_gemm.py | 14 ++++++++++++-- test/cutedsl/test_sm100_gemm.py | 19 +++++++++++++++++++ 2 files changed, 31 insertions(+), 2 deletions(-) diff --git a/src/liger_kernel/ops/cutedsl/ops/_sm100_gemm.py b/src/liger_kernel/ops/cutedsl/ops/_sm100_gemm.py index eeddb665a..0105954c8 100644 --- a/src/liger_kernel/ops/cutedsl/ops/_sm100_gemm.py +++ b/src/liger_kernel/ops/cutedsl/ops/_sm100_gemm.py @@ -878,12 +878,22 @@ def _validate_epilogue_inputs(a, b, out): ) +def _select_epilogue_config(a): + """Select measured SM100 scheduling knobs while retaining the two-CTA kernel.""" + m_tiles = (a.shape[0] + _CTA_M - 1) // _CTA_M + if a.shape[0] >= 4096 and m_tiles % 2 == 0: + if a.shape[1] in (256, 2048): + return 4, 2 + if a.shape[1] in (512, 1024): + return 6, 2 + return _NUM_AB_STAGES, 1 + + def _run_epilogue_gemm(a, b, out, epilogue): _validate_epilogue_inputs(a, b, out) epilogue_key = _validate_epilogue_callback(epilogue) current_stream = _current_stream(a.device) - num_ab_stages = _NUM_AB_STAGES - swizzle_size = 1 + num_ab_stages, swizzle_size = _select_epilogue_config(a) use_tma_output = out.stride(0) * out.element_size() % 16 == 0 max_active_clusters = _max_active_clusters(a.device.index) diff --git a/test/cutedsl/test_sm100_gemm.py b/test/cutedsl/test_sm100_gemm.py index 68e256fd4..6d7e64cf1 100644 --- a/test/cutedsl/test_sm100_gemm.py +++ b/test/cutedsl/test_sm100_gemm.py @@ -44,6 +44,25 @@ def test_sm100_gemm_rejects_non_cute_epilogue(): run_epilogue_gemm(x, weight, out, lambda accumulator, output: None) +@pytest.mark.parametrize( + ("shape", "expected"), + [ + ((128256, 256), (4, 2)), + ((128256, 512), (6, 2)), + ((128256, 1024), (6, 2)), + ((128256, 2048), (4, 2)), + ((512, 4096), (6, 1)), + ((1024, 128256), (6, 1)), + ((128128, 512), (6, 1)), + ], +) +def test_sm100_gemm_selects_measured_config(shape, expected): + import liger_kernel.ops.cutedsl.ops._sm100_gemm as gemm + + tensor = torch.empty(shape, device="meta") + assert gemm._select_epilogue_config(tensor) == expected + + def test_sm100_gemm_custom_epilogue_guards_noncurrent_device(monkeypatch): import liger_kernel.ops.cutedsl.ops._sm100_gemm as gemm From a54c2cb1715cd4d0ea1b238b235a9e9efde3983a Mon Sep 17 00:00:00 2001 From: Justin Hu Date: Fri, 14 Aug 2026 16:08:18 +0000 Subject: [PATCH 2/2] Remove SM100 GEMM dispatch test Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- test/cutedsl/test_sm100_gemm.py | 19 ------------------- 1 file changed, 19 deletions(-) diff --git a/test/cutedsl/test_sm100_gemm.py b/test/cutedsl/test_sm100_gemm.py index 6d7e64cf1..68e256fd4 100644 --- a/test/cutedsl/test_sm100_gemm.py +++ b/test/cutedsl/test_sm100_gemm.py @@ -44,25 +44,6 @@ def test_sm100_gemm_rejects_non_cute_epilogue(): run_epilogue_gemm(x, weight, out, lambda accumulator, output: None) -@pytest.mark.parametrize( - ("shape", "expected"), - [ - ((128256, 256), (4, 2)), - ((128256, 512), (6, 2)), - ((128256, 1024), (6, 2)), - ((128256, 2048), (4, 2)), - ((512, 4096), (6, 1)), - ((1024, 128256), (6, 1)), - ((128128, 512), (6, 1)), - ], -) -def test_sm100_gemm_selects_measured_config(shape, expected): - import liger_kernel.ops.cutedsl.ops._sm100_gemm as gemm - - tensor = torch.empty(shape, device="meta") - assert gemm._select_epilogue_config(tensor) == expected - - def test_sm100_gemm_custom_epilogue_guards_noncurrent_device(monkeypatch): import liger_kernel.ops.cutedsl.ops._sm100_gemm as gemm