[Kernel] Route large SM90 row/column-scaled FP8 GEMMs to Torch (#34318)

This commit is contained in:
Gregory Leleytner
2026-08-29 07:35:37 +08:00
committed by GitHub
parent 96a4dcdde8
commit e1b3bba3cc
3 changed files with 167 additions and 19 deletions
@@ -25,7 +25,7 @@ EXPECTED = {
"activation.relu2": {"jit", "torch", "torch_compile"},
"layernorm.rmsnorm": {"aot", "jit", "aiter", "torch_npu", "torch", "torch_compile"},
"layernorm.gemma_rmsnorm": {"aot", "jit", "torch_npu", "torch", "torch_compile"},
"gemm.fp8_scaled_mm": {"aot"},
"gemm.fp8_scaled_mm": {"aot", "torch", "torch_compile"},
"moe.moe_align_block_size": {"aot", "jit"},
"quantization.nvfp4_gemm_swiglu_nvfp4_quant": {"cute_dsl"},
"kvcache.reshape_and_cache_flash": {"triton"},
@@ -84,7 +84,20 @@ def test_sparse_linear_attention_registry_targets_forward_kernel():
def test_single_backend_resolves_without_backend():
assert K.select_kernel("gemm.fp8_scaled_mm").backend is KernelBackend.AOT
assert (
K.select_kernel("kvcache.reshape_and_cache_flash").backend
is KernelBackend.TRITON
)
def test_fp8_scaled_mm_requires_explicit_registry_backend(monkeypatch):
monkeypatch.setattr(sel, "_platform", lambda: _SM90)
with pytest.raises(ValueError, match="multiple backends"):
K.select_kernel("gemm.fp8_scaled_mm")
assert (
K.select_kernel("gemm.fp8_scaled_mm", backend=KernelBackend.AOT).backend
is KernelBackend.AOT
)
def test_unknown_op_or_backend_raises():