[Kernel] Drop the vendored dense BF16 GEMM port in favor of FlashInfer 0.6.18 (#38124)

Co-authored-by: Mohammad Angkad <mohammad.angkad@radixark.ai>
This commit is contained in:
Mohammad Miadh Angkad
2026-09-07 12:59:33 -07:00
committed by GitHub
co-authored by Mohammad Angkad
parent bf68369a18
commit a711785475
7 changed files with 44 additions and 1473 deletions
@@ -1,367 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES.
# Vendored from flashinfer-ai/flashinfer@629147317d4149a12e53bcef27808bac380c283f.
"""Register-prefetch BF16 GEMM for low-M, long-K decode shapes.
The kernel keeps a complete output dot product inside one CTA and reuses each
prefetched B value across several public-M rows. It is intentionally a
separate autotuner runner from the Blackwell tensor-core split-K kernel: the
two algorithms have different useful shape regions and tactic spaces.
"""
from __future__ import annotations
import dataclasses
import functools
import cuda.bindings.driver as _cuda
import cutlass
import cutlass.cute as cute
import torch as _torch
from cutlass import const_expr
from cutlass.cute import experimental as cute_ext
from cutlass.cute.runtime import from_dlpack
_VECTOR_WIDTH = 8
_SUPPORTED_BLOCK_SIZES = (32, 64, 96, 128, 192, 256, 384)
_SUPPORTED_OUTPUTS_PER_BLOCK = (1, 2, 4)
_MAX_M = 32
_COMPILE_OPTIONS = "--ptxas-options -maxrregcount=64"
@dataclasses.dataclass(frozen=True)
class DirectTactic:
"""One direct-kernel specialization."""
block_size: int
outputs_per_block: int
rows_per_block: int
def _default_rows_per_block(m: int) -> int:
if m <= 8:
return m
return next(rows for rows in (8, 4, 2, 1) if m % rows == 0)
def validate_tactic(tactic: DirectTactic, m: int, n: int, k: int) -> None:
"""Reject a direct tactic that cannot serve ``(m, n, k)``."""
if tactic.block_size not in _SUPPORTED_BLOCK_SIZES:
raise ValueError(f"unsupported block_size={tactic.block_size}")
if tactic.outputs_per_block not in _SUPPORTED_OUTPUTS_PER_BLOCK:
raise ValueError(f"unsupported outputs_per_block={tactic.outputs_per_block}")
if not 1 <= m <= _MAX_M:
raise ValueError(f"direct GEMM requires 1 <= M <= {_MAX_M}, got {m}")
if not 1 <= tactic.rows_per_block <= m or m % tactic.rows_per_block:
raise ValueError(f"rows_per_block={tactic.rows_per_block} must divide M={m}")
if n <= 0 or n % tactic.outputs_per_block:
raise ValueError(
f"N={n} must be divisible by outputs_per_block={tactic.outputs_per_block}"
)
k_tile = tactic.block_size * _VECTOR_WIDTH
if k <= 0 or k % k_tile:
raise ValueError(f"K={k} must be divisible by {k_tile}")
def default_tactic(m: int, n: int, k: int) -> DirectTactic:
"""Choose the measured register-prefetch fallback tactic."""
block_size = next(
(
block
for block in (256, 192, 128, 96, 64, 32)
if k % (block * _VECTOR_WIDTH) == 0
),
None,
)
if block_size is None:
raise ValueError("direct GEMM requires a supported 16-byte K tiling")
outputs_per_block = next(outputs for outputs in (2, 1) if n % outputs == 0)
tactic = DirectTactic(
block_size,
outputs_per_block,
_default_rows_per_block(m),
)
validate_tactic(tactic, m, n, k)
return tactic
def autotune_tactics(m: int, n: int, k: int) -> list[DirectTactic]:
"""Enumerate the compact tactic space used by FlashInfer autotuning.
Block sizes cover every configuration exercised in the H100/B200 sweep;
output grouping spans the measured 1/2/4-column choices. Row tiling stays
at the occupancy-oriented default to keep JIT cost bounded.
"""
try:
default = default_tactic(m, n, k)
except ValueError:
return []
tactics = [default]
for block_size in _SUPPORTED_BLOCK_SIZES:
for outputs_per_block in _SUPPORTED_OUTPUTS_PER_BLOCK:
tactic = DirectTactic(
block_size,
outputs_per_block,
default.rows_per_block,
)
try:
validate_tactic(tactic, m, n, k)
except ValueError:
continue
tactics.append(tactic)
return list(dict.fromkeys(tactics))
def prefer_direct_bf16_gemm_sm100(m: int, n: int, k: int) -> bool:
"""Return the conservative B200 no-autotune crossover heuristic.
The three bands are a compact fit to a warm/cold sweep over M=1..16,24,32,
18 N values, and 11 K values. This is deliberately not a blanket rule for
K=8192: direct wins only where public M and N leave the tensor-core path
with too little independent output work.
"""
return k == 8192 and (
(m == 1 and n <= 4608) or (m <= 4 and n <= 512) or (m <= 8 and n <= 256)
)
class DirectDenseGemmKernel:
"""K-specialized direct GEMM with whole-mainloop vector prefetch."""
def __init__(
self,
*,
element_type,
num_rows: int,
k_extent: int,
tactic: DirectTactic,
use_pdl: bool,
) -> None:
validate_tactic(tactic, num_rows, tactic.outputs_per_block, k_extent)
self.element_type = element_type
self.num_rows = num_rows
self.rows_per_block = tactic.rows_per_block
self.k_extent = k_extent
self.block_size = tactic.block_size
self.outputs_per_block = tactic.outputs_per_block
self.vector_width = _VECTOR_WIDTH
self.use_pdl = use_pdl
self.num_warps = tactic.block_size // cute.arch.WARP_SIZE
self.num_k_tiles = k_extent // (tactic.block_size * _VECTOR_WIDTH)
@cute.jit
def __call__(
self,
gA: cute.Tensor,
gB: cute.Tensor,
gC: cute.Tensor,
stream: _cuda.CUstream,
) -> None:
n = cute.size(gB, mode=[0])
copy_a = cute.make_copy_atom(
cute.nvgpu.CopyG2ROp(),
self.element_type,
num_bits_per_copy=self.vector_width * self.element_type.width,
load_cache_mode=cute.nvgpu.LoadCacheMode.ALWAYS,
)
copy_b = cute.make_copy_atom(
cute.nvgpu.CopyG2ROp(),
self.element_type,
num_bits_per_copy=self.vector_width * self.element_type.width,
load_cache_mode=cute.nvgpu.LoadCacheMode.STREAMING,
)
self.kernel(gA, gB, gC, copy_a, copy_b).launch(
grid=[
cute.ceil_div(n, self.outputs_per_block),
self.num_rows // self.rows_per_block,
1,
],
block=[self.block_size, 1, 1],
smem=self.rows_per_block * self.outputs_per_block * self.num_warps * 4,
stream=stream,
use_pdl=self.use_pdl,
min_blocks_per_mp=1,
)
@cute.kernel
def kernel(
self,
gA: cute.Tensor,
gB: cute.Tensor,
gC: cute.Tensor,
copy_a: cute.CopyAtom,
copy_b: cute.CopyAtom,
) -> None:
tidx, _, _ = cute.arch.thread_idx()
block_idx, block_m, _ = cute.arch.block_idx()
warp_idx = cute.arch.warp_idx()
num_rows: cutlass.Constexpr = self.rows_per_block
outputs_per_block: cutlass.Constexpr = self.outputs_per_block
vector_width: cutlass.Constexpr = self.vector_width
block_size: cutlass.Constexpr = self.block_size
num_warps: cutlass.Constexpr = self.num_warps
num_k_tiles: cutlass.Constexpr = self.num_k_tiles
acc = cute.make_rmem_tensor(
cute.make_layout(
(num_rows, outputs_per_block), stride=(outputs_per_block, 1)
),
cutlass.Float32,
)
acc.fill(0.0)
if const_expr(self.use_pdl):
cute.arch.griddepcontrol_wait()
n_base = block_idx * outputs_per_block
m_base = block_m * num_rows
gA_vec = cute.logical_divide(gA, (None, vector_width))
gB_vec = cute.logical_divide(gB, (None, vector_width))
tA_all = cute.logical_divide(gA_vec, (None, (None, block_size)))
tB_all = cute.logical_divide(gB_vec, (None, (None, block_size)))
tA = tA_all[None, (None, (tidx, None))]
b_regs = cute.make_rmem_tensor(
cute.make_layout(
(outputs_per_block, num_k_tiles, vector_width),
stride=(num_k_tiles * vector_width, vector_width, 1),
),
self.element_type,
)
for ni in cutlass.range_constexpr(outputs_per_block):
tB = tB_all[n_base + ni, (None, (tidx, None))]
for k_tile in cutlass.range_constexpr(num_k_tiles):
cute.copy(copy_b, tB[None, k_tile], b_regs[ni, k_tile, None])
a_regs = cute.make_rmem_tensor(
cute.make_layout((num_k_tiles, vector_width), stride=(vector_width, 1)),
self.element_type,
)
for mi in cutlass.range_constexpr(num_rows):
for k_tile in cutlass.range_constexpr(num_k_tiles):
cute.copy(
copy_a,
tA[m_base + mi, None, k_tile],
a_regs[k_tile, None],
)
for k_tile in cutlass.range_constexpr(num_k_tiles):
for vi in cutlass.range_constexpr(vector_width):
a_value = a_regs[k_tile, vi].to(cutlass.Float32)
for ni in cutlass.range_constexpr(outputs_per_block):
acc[mi, ni] = acc[mi, ni] + a_value * b_regs[ni, k_tile, vi].to(
cutlass.Float32
)
for mi in cutlass.range_constexpr(num_rows):
for ni in cutlass.range_constexpr(outputs_per_block):
acc[mi, ni] = cute.arch.warp_reduction_sum(acc[mi, ni])
smem_layout = cute.make_layout(
(num_rows, outputs_per_block, num_warps),
stride=(outputs_per_block * num_warps, num_warps, 1),
)
smem = cutlass.utils.SmemAllocator()
partials = smem.allocate_tensor(cutlass.Float32, smem_layout, byte_alignment=16)
with cute.arch.elect_one():
for mi in cutlass.range_constexpr(num_rows):
for ni in cutlass.range_constexpr(outputs_per_block):
partials[mi, ni, warp_idx] = acc[mi, ni]
cute.arch.sync_threads()
if tidx == 0:
for mi in cutlass.range_constexpr(num_rows):
for ni in cutlass.range_constexpr(outputs_per_block):
total = cutlass.Float32(0.0)
for warp in cutlass.range_constexpr(num_warps):
total = total + partials[mi, ni, warp]
gC[m_base + mi, n_base + ni] = total.to(self.element_type)
if const_expr(self.use_pdl):
cute.arch.griddepcontrol_launch_dependents()
def _from_dlpack_static(tensor: _torch.Tensor):
# K is specialized and the row stride must retain its 16-byte divisibility
# for the verifier to accept vectorized G2R copies.
return from_dlpack(tensor, assumed_align=32)
def _make_compile_repr_tensors(dtype, m: int, n: int, k: int):
return tuple(
_from_dlpack_static(tensor)
for tensor in (
_torch.empty((m, k), dtype=dtype, device="cuda"),
_torch.empty((n, k), dtype=dtype, device="cuda"),
_torch.empty((m, n), dtype=dtype, device="cuda"),
)
)
@functools.cache
def _get_compiled_direct_kernel(
dtype,
m: int,
n: int,
k: int,
tactic: DirectTactic,
use_pdl: bool,
):
if dtype != _torch.bfloat16:
raise ValueError(f"direct GEMM supports BF16; got {dtype}")
kernel = DirectDenseGemmKernel(
element_type=cutlass.BFloat16,
num_rows=m,
k_extent=k,
tactic=tactic,
use_pdl=use_pdl,
)
tensors = _make_compile_repr_tensors(dtype, m, n, k)
stream = _cuda.CUstream(_torch.cuda.current_stream().cuda_stream)
return cute_ext.compile(kernel, *tensors, stream, options=_COMPILE_OPTIONS)
def _validate_runtime_tensors(a, b, out, tactic: DirectTactic):
if any(not isinstance(tensor, _torch.Tensor) for tensor in (a, b, out)):
raise ValueError("a, b, and out must be torch tensors")
if a.ndim != 2 or b.ndim != 2 or out.ndim != 2:
raise ValueError("direct GEMM accepts only 2D tensors")
if a.device.type != "cuda" or b.device != a.device or out.device != a.device:
raise ValueError("a, b, and out must be on the same CUDA device")
if a.dtype != _torch.bfloat16 or b.dtype != a.dtype or out.dtype != a.dtype:
raise ValueError("a, b, and out must share BF16 dtype")
if not a.is_contiguous() or not b.T.is_contiguous() or not out.is_contiguous():
raise ValueError("direct GEMM requires row-major A/out and column-major B")
if any(tensor.data_ptr() % 32 for tensor in (a, b, out)):
raise ValueError("a, b, and out must be 32-byte aligned")
m, k = a.shape
if b.shape[0] != k:
raise ValueError(
f"incompatible shapes: a is {tuple(a.shape)}, b is {tuple(b.shape)}"
)
n = b.shape[1]
if out.shape != (m, n):
raise ValueError(f"out must have shape {(m, n)}, got {tuple(out.shape)}")
validate_tactic(tactic, m, n, k)
return m, n, k
def run_direct_dense(a, b, out, pdl: bool, tactic: DirectTactic):
"""Run direct ``A[M,K] @ B[K,N]`` with the ``mm_bf16`` layouts."""
m, n, k = _validate_runtime_tensors(a, b, out, tactic)
compiled = _get_compiled_direct_kernel(a.dtype, m, n, k, tactic, pdl)
tensors = tuple(_from_dlpack_static(tensor) for tensor in (a, b.T, out))
stream = _cuda.CUstream(_torch.cuda.current_stream(a.device).cuda_stream)
compiled(*tensors, stream)
return out
__all__ = [
"DirectTactic",
"autotune_tactics",
"default_tactic",
"prefer_direct_bf16_gemm_sm100",
"run_direct_dense",
]
File diff suppressed because it is too large Load Diff
+2 -2
View File
@@ -178,9 +178,9 @@ class ExecKernel:
bf16_gemm_backend: A[
str,
Arg(
help="Choose the backend for unquantized BF16 GEMM operations. Options: 'auto' (default; selects 'cutedsl' on SM10x GPUs, except deterministic inference selects 'torch'; otherwise uses cuBLAS via torch.nn.functional.linear), 'cutedsl' (SGLang JIT CuTe DSL TGV BF16 GEMM on SM10x; dispatches between the allowlisted low-M Split-K kernel, the CuTe DSL kernel, and cuBLAS; set SGLANG_ENABLE_BF16_SPLITK_GEMM=0 to disable Split-K), 'flashinfer_pr4266' (legacy compatibility alias for the optimized CuTe DSL path), 'gemv', 'torch' (always uses cuBLAS via torch.nn.functional.linear).",
help="Choose the backend for unquantized BF16 GEMM operations. Options: 'auto' (default; selects 'cutedsl' on SM10x GPUs, except deterministic inference selects 'torch'; otherwise uses cuBLAS via torch.nn.functional.linear), 'cutedsl' (SGLang JIT CuTe DSL TGV BF16 GEMM on SM10x; dispatches between the allowlisted low-M Split-K kernel, the CuTe DSL kernel, and cuBLAS; set SGLANG_ENABLE_BF16_SPLITK_GEMM=0 to disable Split-K), 'gemv', 'torch' (always uses cuBLAS via torch.nn.functional.linear).",
cli_name="--bf16-gemm-backend",
choices=["auto", "cutedsl", "flashinfer_pr4266", "gemv", "torch"],
choices=["auto", "cutedsl", "gemv", "torch"],
),
] = "auto"
dsa_prefill_backend: A[
-2
View File
@@ -1794,8 +1794,6 @@ _DEPRECATED_ENVS: Dict[str, _DeprecatedEnv] = {
"SGLANG_OPT_SWA_EVICT_DROP_PAGE_MARGIN": _DeprecatedEnv(),
# sconv-family kernels always use the CUDA-JIT ports when supported; no toggle.
"SGLANG_OPT_USE_CUDA_SCONV": _DeprecatedEnv(),
# The direct dense BF16 GEMM source is vendored in-tree.
"SGLANG_FLASHINFER_PR4266_SOURCE": _DeprecatedEnv(),
# DSV4 compressor V2 is always used.
"SGLANG_OPT_USE_COMPRESSOR_V2": _DeprecatedEnv(),
"SGLANG_ENABLE_HICACHE_BUFFER_ANCHOR_LOCK": _DeprecatedEnv(
@@ -76,7 +76,6 @@ if _use_aiter:
class Bf16GemmBackend(Enum):
AUTO = "auto"
CUTEDSL = "cutedsl"
FLASHINFER_PR4266 = "flashinfer_pr4266"
GEMV = "gemv"
TORCH = "torch"
@@ -89,28 +88,22 @@ class Bf16GemmBackend(Enum):
def is_gemv(self) -> bool:
return self == Bf16GemmBackend.GEMV
def is_flashinfer_pr4266(self) -> bool:
return self == Bf16GemmBackend.FLASHINFER_PR4266
def is_optimized(self) -> bool:
return self.is_cutedsl() or self.is_flashinfer_pr4266()
_BF16_GEMM_BACKEND: Optional[Bf16GemmBackend] = None
_cutedsl_bf16_gemm = None
_use_cutedsl_bf16_gemm = None
_hopper_bf16_gemv = None
_use_hopper_bf16_gemv = None
_flashinfer_pr4266_splitk_tactic = None
_flashinfer_pr4266_run_splitk_dense = None
_flashinfer_pr4266_direct_default_tactic = None
_flashinfer_pr4266_prefer_direct = None
_flashinfer_pr4266_run_direct_dense = None
_splitk_tactic = None
_run_splitk_dense = None
_direct_default_tactic = None
_prefer_direct = None
_run_direct_dense = None
_enable_bf16_splitk_gemm = False
# GB300 TP16 tactics measured under CUDA graph replay with PDL and cold weights.
# Unlisted shapes, including M=64, retain the existing TGV/cuBLAS path.
_FLASHINFER_PR4266_TUNED_TACTICS = {
_BF16_SPLITK_TUNED_TACTICS = {
(1, 256, 8192): (64, 8, 4, 11),
(2, 256, 8192): (64, 8, 4, 11),
(4, 256, 8192): (64, 8, 4, 11),
@@ -142,23 +135,23 @@ _FLASHINFER_PR4266_TUNED_TACTICS = {
}
def use_flashinfer_pr4266_bf16_gemm(m: int, n: int, k: int) -> bool:
return (m, n, k) in _FLASHINFER_PR4266_TUNED_TACTICS
def use_bf16_splitk_gemm(m: int, n: int, k: int) -> bool:
return (m, n, k) in _BF16_SPLITK_TUNED_TACTICS
def should_enable_bf16_splitk_gemm(backend: Bf16GemmBackend) -> bool:
"""Return whether the optional Split-K path should be initialized."""
return backend.is_optimized() and envs.SGLANG_ENABLE_BF16_SPLITK_GEMM.get()
return backend.is_cutedsl() and envs.SGLANG_ENABLE_BF16_SPLITK_GEMM.get()
def initialize_bf16_gemm_config() -> None:
global _BF16_GEMM_BACKEND
global _cutedsl_bf16_gemm, _use_cutedsl_bf16_gemm
global _flashinfer_pr4266_splitk_tactic
global _flashinfer_pr4266_run_splitk_dense
global _flashinfer_pr4266_direct_default_tactic
global _flashinfer_pr4266_prefer_direct
global _flashinfer_pr4266_run_direct_dense
global _splitk_tactic
global _run_splitk_dense
global _direct_default_tactic
global _prefer_direct
global _run_direct_dense
global _enable_bf16_splitk_gemm
backend_str = get_exec().kernel.bf16_gemm_backend
@@ -183,7 +176,7 @@ def initialize_bf16_gemm_config() -> None:
_hopper_bf16_gemv = hopper_bf16_gemv
_use_hopper_bf16_gemv = use_hopper_bf16_gemv
elif backend.is_optimized():
elif backend.is_cutedsl():
if get_exec().deterministic.enable_deterministic_inference:
raise ValueError(
"--bf16-gemm-backend cutedsl is batch-size dependent and cannot "
@@ -204,21 +197,21 @@ def initialize_bf16_gemm_config() -> None:
_enable_bf16_splitk_gemm = False
if should_enable_bf16_splitk_gemm(backend):
from sglang.kernels.ops.gemm.flashinfer_pr4266_dense_bf16_gemm_sm100_direct import (
from flashinfer.gemm.kernels.dense_bf16_gemm_direct import (
default_tactic,
prefer_direct_bf16_gemm_sm100,
run_direct_dense,
)
from sglang.kernels.ops.gemm.flashinfer_pr4266_dense_bf16_gemm_sm100_splitk import (
from flashinfer.gemm.kernels.dense_bf16_gemm_sm100_splitk import (
SplitKTactic,
run_splitk_dense,
)
_flashinfer_pr4266_splitk_tactic = SplitKTactic
_flashinfer_pr4266_run_splitk_dense = run_splitk_dense
_flashinfer_pr4266_direct_default_tactic = default_tactic
_flashinfer_pr4266_prefer_direct = prefer_direct_bf16_gemm_sm100
_flashinfer_pr4266_run_direct_dense = run_direct_dense
_splitk_tactic = SplitKTactic
_run_splitk_dense = run_splitk_dense
_direct_default_tactic = default_tactic
_prefer_direct = prefer_direct_bf16_gemm_sm100
_run_direct_dense = run_direct_dense
_enable_bf16_splitk_gemm = True
_BF16_GEMM_BACKEND = backend
@@ -230,20 +223,18 @@ def _bf16_gemm_dispatch_fake(
return x.new_empty((*x.shape[:-1], weight.shape[0]))
def _flashinfer_pr4266_bf16_gemm(
def _bf16_splitk_gemm(
x: torch.Tensor, weight: torch.Tensor, bias: Optional[torch.Tensor]
) -> torch.Tensor:
x_2d = x.view(-1, x.shape[-1])
out = torch.empty((x_2d.shape[0], weight.shape[0]), dtype=x.dtype, device=x.device)
m, n, k = x_2d.shape[0], weight.shape[0], weight.shape[1]
if bias is None and _flashinfer_pr4266_prefer_direct(m, n, k):
tactic = _flashinfer_pr4266_direct_default_tactic(m, n, k)
_flashinfer_pr4266_run_direct_dense(x_2d, weight.T, out, True, tactic)
if bias is None and _prefer_direct(m, n, k):
tactic = _direct_default_tactic(m, n, k)
_run_direct_dense(x_2d, weight.T, out, True, tactic)
else:
tactic = _flashinfer_pr4266_splitk_tactic(
*_FLASHINFER_PR4266_TUNED_TACTICS[(m, n, k)]
)
_flashinfer_pr4266_run_splitk_dense(
tactic = _splitk_tactic(*_BF16_SPLITK_TUNED_TACTICS[(m, n, k)])
_run_splitk_dense(
x_2d,
weight.T,
bias,
@@ -261,10 +252,10 @@ def _bf16_gemm_dispatch_impl(
addend: Optional[torch.Tensor] = None,
) -> torch.Tensor:
m = x.numel() // x.shape[-1]
if _enable_bf16_splitk_gemm and use_flashinfer_pr4266_bf16_gemm(
if _enable_bf16_splitk_gemm and use_bf16_splitk_gemm(
m, weight.shape[0], weight.shape[1]
):
output = _flashinfer_pr4266_bf16_gemm(x, weight, bias)
output = _bf16_splitk_gemm(x, weight, bias)
elif (
_use_hopper_bf16_gemv is not None
and bias is None
@@ -423,7 +414,7 @@ class UnquantizedLinearMethod(LinearMethodBase):
return tgemm.mm(x, layer.weight, bias, otype=x.dtype)
elif (
get_bf16_gemm_backend().is_optimized()
get_bf16_gemm_backend().is_cutedsl()
and x.is_cuda
and x.dtype == torch.bfloat16
and layer.weight.dtype == torch.bfloat16
@@ -2,25 +2,25 @@ import pytest
from sglang.srt.environ import envs
from sglang.srt.layers.quantization.unquant import (
_FLASHINFER_PR4266_TUNED_TACTICS,
_BF16_SPLITK_TUNED_TACTICS,
Bf16GemmBackend,
should_enable_bf16_splitk_gemm,
use_flashinfer_pr4266_bf16_gemm,
use_bf16_splitk_gemm,
)
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=11, suite="base-a-test-cpu")
@pytest.mark.parametrize("m,n,k", _FLASHINFER_PR4266_TUNED_TACTICS)
def test_flashinfer_pr4266_selects_tuned_oakhaven_shape(m: int, n: int, k: int):
assert use_flashinfer_pr4266_bf16_gemm(m, n, k)
@pytest.mark.parametrize("m,n,k", _BF16_SPLITK_TUNED_TACTICS)
def test_splitk_selects_tuned_oakhaven_shape(m: int, n: int, k: int):
assert use_bf16_splitk_gemm(m, n, k)
@pytest.mark.parametrize("m", [0, 33, 64])
@pytest.mark.parametrize("n,k", [(256, 8192), (512, 8192), (2304, 8192), (2560, 8192)])
def test_flashinfer_pr4266_keeps_large_m_on_existing_path(m: int, n: int, k: int):
assert not use_flashinfer_pr4266_bf16_gemm(m, n, k)
def test_splitk_keeps_large_m_on_existing_path(m: int, n: int, k: int):
assert not use_bf16_splitk_gemm(m, n, k)
@pytest.mark.parametrize(
@@ -32,12 +32,8 @@ def test_flashinfer_pr4266_keeps_large_m_on_existing_path(m: int, n: int, k: int
(32, 4096, 8192),
],
)
def test_flashinfer_pr4266_rejects_unmeasured_shapes(shape: tuple[int, int, int]):
assert not use_flashinfer_pr4266_bf16_gemm(*shape)
def test_flashinfer_pr4266_backend_is_explicit():
assert Bf16GemmBackend.FLASHINFER_PR4266.value == "flashinfer_pr4266"
def test_splitk_rejects_unmeasured_shapes(shape: tuple[int, int, int]):
assert not use_bf16_splitk_gemm(*shape)
def test_bf16_splitk_is_enabled_by_default():
@@ -168,15 +168,11 @@ class TestApplyWithAddend(CustomTestCase):
)
elif route == "splitk":
enter(patch.object(unquant, "_enable_bf16_splitk_gemm", True))
enter(
patch.object(
unquant, "use_flashinfer_pr4266_bf16_gemm", lambda *a: True
)
)
enter(patch.object(unquant, "use_bf16_splitk_gemm", lambda *a: True))
enter(
patch.object(
unquant,
"_flashinfer_pr4266_bf16_gemm",
"_bf16_splitk_gemm",
_fake_kernel(kernel_calls, route),
)
)