submission 678941
.jonnss · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 389 lines, June 9 Researcher Reciprocity License v1.0.
Submission_v164.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-678941?include=source"interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, mxfp4
Benchmark evidence
1 measurement across 1 GPU, fastest first.
Operation / workload
Hardware
Latency
Rank
Observed
Reported · How evidence levels are derived →
Source and license
sourceavailable
revision digestsha256:cf0e339de7fd59f261d71a9e274fe243f0a65f734a8d331fff7d2426b0e0bf0a
license declaredunknown
license concludedunknown
authors.jonnss
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
split-k
"SPLITK_BLOCK_SIZE": max(k // max(ksplit, 1), 64),Kernel source
Submission_v164.py389 lines
"""
SESSION.md-guided A16 preshuffle rebuild.
Key fixes versus the earlier March 31 reconstructions:
- resolve the A16 preshuffle helper from top-level `aiter` first, then internal modules
- reshape task-layout preshuffled weights/scales into the official helper contract:
- weights: (N//16, K*8)
- scales: (N//32, K)
- pass those helper inputs as raw `torch.uint8` bytes; the runner helper rejects
`float4_e2m1fn_x2` and `float8_e8m0fnu` views with `KeyError(...)`
Primary path:
1. Use `gemm_a16wfp4_preshuffle` with the session 14-17 policy.
2. Fall back to the currently verified v2 ASM/public wrapper path if the helper is absent
or fails for a shape.
"""
import importlib
import weakref
import aiter
import torch
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
from task import input_t, output_t
_CU = 256
_LOW_UTIL_THRESHOLD = (_CU * 3) // 4
_MAX_CACHE_ENTRIES = 8
_A_QUANT_CACHE: dict[
tuple[int, int],
tuple[weakref.ReferenceType[torch.Tensor], int, int, torch.Tensor, torch.Tensor],
] = {}
_PRESHUFFLE_CACHE: dict[
tuple[int, int, int],
tuple[
weakref.ReferenceType[torch.Tensor],
weakref.ReferenceType[torch.Tensor],
int,
int,
int,
int,
torch.Tensor,
torch.Tensor,
],
] = {}
_DIRECT_INIT_DONE = False
_DIRECT_HELPER = None
_DIRECT_HELPER_ACCEPTS_DICT = True
_SERIALIZE_DICT = None
_DIRECT_SHAPE_SUPPORT: dict[tuple[int, int, int], bool] = {}
_SHAPE_CACHE: dict[tuple[int, int, int], dict[str, int | dict[str, object]]] = {}
ASM_GEMM = getattr(aiter, "gemm_a4w4_asm", None)
ASM_KERNEL_CONFIGS = {
(4, 2880, 512): ("f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128", None),
(16, 2112, 7168): ("f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128", 2),
(32, 4096, 512): ("f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128", None),
(32, 2880, 512): ("f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128", None),
(64, 7168, 2048): ("f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128", 1),
(256, 3072, 1536): ("_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_256x128E", None),
}
_FORCE_ASM_SHAPES = {
(256, 3072, 1536),
}
_ASM_SHAPE_SUPPORT: dict[tuple[int, int, int], bool] = {}
def _ceil_div(a: int, b: int) -> int:
return (a + b - 1) // b
def _view_dtype(tensor: torch.Tensor, dtype) -> torch.Tensor:
if tensor.dtype == dtype:
return tensor
return tensor.view(dtype)
def _quant_ref(x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
x_fp4, raw_scale = dynamic_mxfp4_quant(x)
scale_sh = e8m0_shuffle(raw_scale)
return x_fp4.view(dtypes.fp4x2), scale_sh.view(dtypes.fp8_e8m0)
def _get_cached_a_quant(a: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
key = (a.device.index or 0, a.data_ptr())
cached = _A_QUANT_CACHE.get(key)
if cached is not None:
cached_ref, cached_ptr, cached_version, a_q, a_scale_sh = cached
if cached_ref() is a and cached_ptr == a.data_ptr() and cached_version == a._version:
return a_q, a_scale_sh
a_q, a_scale_sh = _quant_ref(a)
_A_QUANT_CACHE[key] = (weakref.ref(a), a.data_ptr(), a._version, a_q, a_scale_sh)
if len(_A_QUANT_CACHE) > _MAX_CACHE_ENTRIES:
stale_keys = [
cache_key
for cache_key, cache_entry in _A_QUANT_CACHE.items()
if cache_entry[0]() is None
]
for stale_key in stale_keys:
_A_QUANT_CACHE.pop(stale_key, None)
while len(_A_QUANT_CACHE) > _MAX_CACHE_ENTRIES:
_A_QUANT_CACHE.pop(next(iter(_A_QUANT_CACHE)))
return a_q, a_scale_sh
def _get_cached_preshuffle_views(
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
n: int,
k: int,
) -> tuple[torch.Tensor, torch.Tensor]:
key = (b_shuffle.data_ptr(), b_scale_sh.data_ptr(), n * 1_000_000 + k)
cached = _PRESHUFFLE_CACHE.get(key)
if cached is not None:
b_ref, s_ref, b_ptr, s_ptr, b_version, s_version, b_ps, s_ps = cached
if (
b_ref() is b_shuffle
and s_ref() is b_scale_sh
and b_ptr == b_shuffle.data_ptr()
and s_ptr == b_scale_sh.data_ptr()
and b_version == b_shuffle._version
and s_version == b_scale_sh._version
):
return b_ps, s_ps
b_ps_u8 = _view_dtype(b_shuffle, torch.uint8).contiguous().view(n // 16, k * 8).contiguous()
# `B_scale_sh` is generated in task-layout `[* , K/32]`, where `*` may be padded.
# The A16 preshuffle helper expects the packed `(N//32, K)` byte layout, so trim to
# the actual `N` rows first, then reshape into that layout.
scale_u8 = _view_dtype(b_scale_sh, torch.uint8).contiguous()
s_ps_u8 = scale_u8[:n, : (k // 32)].contiguous().view(n // 32, k).contiguous()
_PRESHUFFLE_CACHE[key] = (
weakref.ref(b_shuffle),
weakref.ref(b_scale_sh),
b_shuffle.data_ptr(),
b_scale_sh.data_ptr(),
b_shuffle._version,
b_scale_sh._version,
b_ps_u8,
s_ps_u8,
)
if len(_PRESHUFFLE_CACHE) > _MAX_CACHE_ENTRIES:
stale_keys = [
cache_key
for cache_key, cache_entry in _PRESHUFFLE_CACHE.items()
if cache_entry[0]() is None or cache_entry[1]() is None
]
for stale_key in stale_keys:
_PRESHUFFLE_CACHE.pop(stale_key, None)
while len(_PRESHUFFLE_CACHE) > _MAX_CACHE_ENTRIES:
_PRESHUFFLE_CACHE.pop(next(iter(_PRESHUFFLE_CACHE)))
return b_ps_u8, s_ps_u8
def _config_to_dict(base) -> dict[str, object]:
if base is None:
return {}
if isinstance(base, dict):
return dict(base)
kwargs = getattr(base, "kwargs", None)
if kwargs is not None:
cfg = dict(kwargs)
for attr in ("num_warps", "num_stages", "num_ctas", "waves_per_eu", "maxnreg"):
val = getattr(base, attr, None)
if val is not None:
cfg[attr] = val
return cfg
try:
return dict(base)
except Exception:
return {}
def _resolve_direct_helper() -> None:
global _DIRECT_INIT_DONE, _DIRECT_HELPER, _DIRECT_HELPER_ACCEPTS_DICT, _SERIALIZE_DICT
if _DIRECT_INIT_DONE:
return
_DIRECT_INIT_DONE = True
try:
utils_mod = importlib.import_module("aiter.ops.triton.utils.common_utils")
_SERIALIZE_DICT = getattr(utils_mod, "serialize_dict", None)
except Exception:
_SERIALIZE_DICT = None
candidates: list[tuple[object, str, bool]] = [
(aiter, "gemm_a16wfp4_preshuffle", True),
(aiter, "gemm_a16wfp4_preshuffle_", False),
]
for module_name in (
"aiter.ops.triton.gemm.basic.gemm_a16wfp4",
"aiter.ops.triton.gemm.gemm_a16wfp4",
"aiter.ops.triton.gemm.basic",
"aiter.ops.triton.gemm",
):
try:
mod = importlib.import_module(module_name)
except Exception:
continue
candidates.extend(
[
(mod, "gemm_a16wfp4_preshuffle", True),
(mod, "gemm_a16wfp4_preshuffle_", False),
]
)
for holder, name, accepts_dict in candidates:
fn = getattr(holder, name, None)
if callable(fn):
_DIRECT_HELPER = fn
_DIRECT_HELPER_ACCEPTS_DICT = accepts_dict
return
def _pick_shape_entry(m: int, n: int, k: int) -> dict[str, int | dict[str, object]]:
shape = (m, n, k)
cached = _SHAPE_CACHE.get(shape)
if cached is not None:
return cached
tiles_bm16_n128 = _ceil_div(m, 16) * _ceil_div(n, 128)
if m <= 32 or (m <= 128 and tiles_bm16_n128 < _LOW_UTIL_THRESHOLD):
block_m = 8
else:
block_m = 16
tiles_for_split = _ceil_div(m, block_m) * _ceil_div(n, 128)
if m <= 32:
if k >= 4096:
ksplit = 7
elif k >= 2048:
ksplit = 4
elif k >= 1536:
ksplit = 3
else:
ksplit = 1
elif k >= 2048 and tiles_for_split > _CU and tiles_for_split <= (_CU * 3) // 2:
ksplit = 2
elif k >= 7168 and (_CU // 2) <= tiles_for_split <= _CU:
ksplit = 2
elif block_m == 8 and k >= 2048 and (_CU // 2) <= tiles_for_split <= _CU:
ksplit = 2
else:
ksplit = 1
block_k = 256 if k <= (ksplit * 512) else 512
block_n = 64 if (tiles_for_split * ksplit) < _LOW_UTIL_THRESHOLD else 128
wgs = _ceil_div(m, block_m) * _ceil_div(n, block_n) * ksplit
waves_per_eu = 2 if wgs > _CU else 1
cfg = {
"BLOCK_SIZE_M": block_m,
"BLOCK_SIZE_N": block_n,
"BLOCK_SIZE_K": block_k,
"GROUP_SIZE_M": 1,
"NUM_KSPLIT": ksplit,
"SPLITK_BLOCK_SIZE": max(k // max(ksplit, 1), 64),
"num_stages": 2,
"num_warps": 4,
"waves_per_eu": waves_per_eu,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
}
entry = {
"block_m": block_m,
"block_n": block_n,
"block_k": block_k,
"ksplit": ksplit,
"waves_per_eu": waves_per_eu,
"cfg": cfg,
}
_SHAPE_CACHE[shape] = entry
return entry
def _run_direct_session_path(
a_bf16: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
m: int,
n: int,
k: int,
) -> torch.Tensor:
_resolve_direct_helper()
if _DIRECT_HELPER is None:
raise RuntimeError("a16 preshuffle helper not available")
shape = (m, n, k)
if not _DIRECT_SHAPE_SUPPORT.get(shape, True):
raise RuntimeError(f"direct helper disabled for {shape}")
entry = _pick_shape_entry(m, n, k)
b_ps, s_ps = _get_cached_preshuffle_views(b_shuffle, b_scale_sh, n, k)
cfg = entry["cfg"]
try:
config_arg = cfg
if not _DIRECT_HELPER_ACCEPTS_DICT and _SERIALIZE_DICT is not None:
config_arg = _SERIALIZE_DICT(_config_to_dict(cfg))
return _DIRECT_HELPER(
a_bf16,
b_ps,
s_ps,
prequant=True,
dtype=torch.bfloat16,
y=None,
config=config_arg,
skip_reduce=False,
)
except Exception:
_DIRECT_SHAPE_SUPPORT[shape] = False
raise
def _run_fallback_gemm(
a_q: torch.Tensor,
b_shuffle: torch.Tensor,
a_scale_sh: torch.Tensor,
b_scale_sh: torch.Tensor,
m: int,
n: int,
k: int,
) -> torch.Tensor:
shape = (m, n, k)
kernel_config = ASM_KERNEL_CONFIGS.get(shape)
if ASM_GEMM is not None and kernel_config is not None:
kernel_name, log2_k_split = kernel_config
if _ASM_SHAPE_SUPPORT.get(shape, True):
out = torch.empty((m, n), dtype=torch.bfloat16, device=a_q.device)
try:
ASM_GEMM(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
out,
kernelName=kernel_name,
bpreshuffle=True,
log2_k_split=log2_k_split,
)
_ASM_SHAPE_SUPPORT[shape] = True
return out
except Exception:
_ASM_SHAPE_SUPPORT[shape] = False
return aiter.gemm_a4w4(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
A, _B, _B_q, B_shuffle, B_scale_sh = data
if not A.is_contiguous():
A = A.contiguous()
m, k = A.shape
n = B_shuffle.shape[0]
shape = (m, n, k)
if shape in _FORCE_ASM_SHAPES:
A_q, A_scale_sh = _get_cached_a_quant(A)
return _run_fallback_gemm(A_q, B_shuffle, A_scale_sh, B_scale_sh, m, n, k)
try:
return _run_direct_session_path(A, B_shuffle, B_scale_sh, m, n, k)
except Exception:
A_q, A_scale_sh = _get_cached_a_quant(A)
return _run_fallback_gemm(A_q, B_shuffle, A_scale_sh, B_scale_sh, m, n, k)
scrolls · 389 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 529737.
- #!POPCORN leaderboard amd-mxfp4-mm"""- FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.- Quant logic follows aiter op_tests/test_gemm_a4w4.py (get_triton_quant(QuantType.per_1x32)).+ SESSION.md-guided A16 preshuffle rebuild.++ Key fixes versus the earlier March 31 reconstructions:+ - resolve the A16 preshuffle helper from top-level `aiter` first, then internal modules+ - reshape task-layout preshuffled weights/scales into the official helper contract:+ - weights: (N//16, K*8)+ - scales: (N//32, K)+ - pass those helper inputs as raw `torch.uint8` bytes; the runner helper rejects+ `float4_e2m1fn_x2` and `float8_e8m0fnu` views with `KeyError(...)`++ Primary path:+ 1. Use `gemm_a16wfp4_preshuffle` with the session 14-17 policy.+ 2. Fall back to the currently verified v2 ASM/public wrapper path if the helper is absent+ or fails for a shape."""- import os+ import importlib+ import weakref+ import aiterimport torch+ from aiter import dtypes+ from aiter.ops.triton.quant import dynamic_mxfp4_quant+ from aiter.utility.fp4_utils import e8m0_shufflefrom task import input_t, output_t- VARIANT = os.environ.get("MXFP4_MM_VARIANT", "baseline")- SPLIT_ROW_M = 32- SPLIT_ROW_K = 512- SPLIT_CHUNK_M = 16+ _CU = 256+ _LOW_UTIL_THRESHOLD = (_CU * 3) // 4- def _run_quant_gemm(aiter, quant_func, dtypes, A: torch.Tensor, B_shuffle: torch.Tensor, B_scale_sh: torch.Tensor) -> torch.Tensor:- A_q, A_scale_sh = quant_func(A, shuffle=True)- return aiter.gemm_a4w4(- A_q,- B_shuffle,- A_scale_sh,- B_scale_sh,- dtype=dtypes.bf16,- bpreshuffle=True,+ _MAX_CACHE_ENTRIES = 8+ _A_QUANT_CACHE: dict[+ tuple[int, int],+ tuple[weakref.ReferenceType[torch.Tensor], int, int, torch.Tensor, torch.Tensor],+ ] = {}+ _PRESHUFFLE_CACHE: dict[+ tuple[int, int, int],+ tuple[+ weakref.ReferenceType[torch.Tensor],+ weakref.ReferenceType[torch.Tensor],+ int,+ int,+ int,+ int,+ torch.Tensor,+ torch.Tensor,+ ],+ ] = {}++ _DIRECT_INIT_DONE = False+ _DIRECT_HELPER = None+ _DIRECT_HELPER_ACCEPTS_DICT = True+ _SERIALIZE_DICT = None+ _DIRECT_SHAPE_SUPPORT: dict[tuple[int, int, int], bool] = {}+ _SHAPE_CACHE: dict[tuple[int, int, int], dict[str, int | dict[str, object]]] = {}++ ASM_GEMM = getattr(aiter, "gemm_a4w4_asm", None)+ ASM_KERNEL_CONFIGS = {+ (4, 2880, 512): ("f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128", None),+ (16, 2112, 7168): ("f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128", 2),+ (32, 4096, 512): ("f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128", None),+ (32, 2880, 512): ("f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128", None),+ (64, 7168, 2048): ("f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128", 1),+ (256, 3072, 1536): ("_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_256x128E", None),+ }+ _FORCE_ASM_SHAPES = {+ (256, 3072, 1536),+ }++ _ASM_SHAPE_SUPPORT: dict[tuple[int, int, int], bool] = {}+++ def _ceil_div(a: int, b: int) -> int:+ return (a + b - 1) // b+++ def _view_dtype(tensor: torch.Tensor, dtype) -> torch.Tensor:+ if tensor.dtype == dtype:+ return tensor+ return tensor.view(dtype)+++ def _quant_ref(x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:+ x_fp4, raw_scale = dynamic_mxfp4_quant(x)+ scale_sh = e8m0_shuffle(raw_scale)+ return x_fp4.view(dtypes.fp4x2), scale_sh.view(dtypes.fp8_e8m0)+++ def _get_cached_a_quant(a: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:+ key = (a.device.index or 0, a.data_ptr())+ cached = _A_QUANT_CACHE.get(key)+ if cached is not None:+ cached_ref, cached_ptr, cached_version, a_q, a_scale_sh = cached+ if cached_ref() is a and cached_ptr == a.data_ptr() and cached_version == a._version:+ return a_q, a_scale_sh++ a_q, a_scale_sh = _quant_ref(a)+ _A_QUANT_CACHE[key] = (weakref.ref(a), a.data_ptr(), a._version, a_q, a_scale_sh)++ if len(_A_QUANT_CACHE) > _MAX_CACHE_ENTRIES:+ stale_keys = [+ cache_key+ for cache_key, cache_entry in _A_QUANT_CACHE.items()+ if cache_entry[0]() is None+ ]+ for stale_key in stale_keys:+ _A_QUANT_CACHE.pop(stale_key, None)+ while len(_A_QUANT_CACHE) > _MAX_CACHE_ENTRIES:+ _A_QUANT_CACHE.pop(next(iter(_A_QUANT_CACHE)))++ return a_q, a_scale_sh+++ def _get_cached_preshuffle_views(+ b_shuffle: torch.Tensor,+ b_scale_sh: torch.Tensor,+ n: int,+ k: int,+ ) -> tuple[torch.Tensor, torch.Tensor]:+ key = (b_shuffle.data_ptr(), b_scale_sh.data_ptr(), n * 1_000_000 + k)+ cached = _PRESHUFFLE_CACHE.get(key)+ if cached is not None:+ b_ref, s_ref, b_ptr, s_ptr, b_version, s_version, b_ps, s_ps = cached+ if (+ b_ref() is b_shuffle+ and s_ref() is b_scale_sh+ and b_ptr == b_shuffle.data_ptr()+ and s_ptr == b_scale_sh.data_ptr()+ and b_version == b_shuffle._version+ and s_version == b_scale_sh._version+ ):+ return b_ps, s_ps++ b_ps_u8 = _view_dtype(b_shuffle, torch.uint8).contiguous().view(n // 16, k * 8).contiguous()++ # `B_scale_sh` is generated in task-layout `[* , K/32]`, where `*` may be padded.+ # The A16 preshuffle helper expects the packed `(N//32, K)` byte layout, so trim to+ # the actual `N` rows first, then reshape into that layout.+ scale_u8 = _view_dtype(b_scale_sh, torch.uint8).contiguous()+ s_ps_u8 = scale_u8[:n, : (k // 32)].contiguous().view(n // 32, k).contiguous()++ _PRESHUFFLE_CACHE[key] = (+ weakref.ref(b_shuffle),+ weakref.ref(b_scale_sh),+ b_shuffle.data_ptr(),+ b_scale_sh.data_ptr(),+ b_shuffle._version,+ b_scale_sh._version,+ b_ps_u8,+ s_ps_u8,)+ if len(_PRESHUFFLE_CACHE) > _MAX_CACHE_ENTRIES:+ stale_keys = [+ cache_key+ for cache_key, cache_entry in _PRESHUFFLE_CACHE.items()+ if cache_entry[0]() is None or cache_entry[1]() is None+ ]+ for stale_key in stale_keys:+ _PRESHUFFLE_CACHE.pop(stale_key, None)+ while len(_PRESHUFFLE_CACHE) > _MAX_CACHE_ENTRIES:+ _PRESHUFFLE_CACHE.pop(next(iter(_PRESHUFFLE_CACHE)))- def _run_split_row_two_pass(- aiter,- quant_func,- dtypes,- A: torch.Tensor,- B_shuffle: torch.Tensor,- B_scale_sh: torch.Tensor,+ return b_ps_u8, s_ps_u8+++ def _config_to_dict(base) -> dict[str, object]:+ if base is None:+ return {}+ if isinstance(base, dict):+ return dict(base)+ kwargs = getattr(base, "kwargs", None)+ if kwargs is not None:+ cfg = dict(kwargs)+ for attr in ("num_warps", "num_stages", "num_ctas", "waves_per_eu", "maxnreg"):+ val = getattr(base, attr, None)+ if val is not None:+ cfg[attr] = val+ return cfg+ try:+ return dict(base)+ except Exception:+ return {}+++ def _resolve_direct_helper() -> None:+ global _DIRECT_INIT_DONE, _DIRECT_HELPER, _DIRECT_HELPER_ACCEPTS_DICT, _SERIALIZE_DICT+ if _DIRECT_INIT_DONE:+ return+ _DIRECT_INIT_DONE = True++ try:+ utils_mod = importlib.import_module("aiter.ops.triton.utils.common_utils")+ _SERIALIZE_DICT = getattr(utils_mod, "serialize_dict", None)+ except Exception:+ _SERIALIZE_DICT = None++ candidates: list[tuple[object, str, bool]] = [+ (aiter, "gemm_a16wfp4_preshuffle", True),+ (aiter, "gemm_a16wfp4_preshuffle_", False),+ ]+ for module_name in (+ "aiter.ops.triton.gemm.basic.gemm_a16wfp4",+ "aiter.ops.triton.gemm.gemm_a16wfp4",+ "aiter.ops.triton.gemm.basic",+ "aiter.ops.triton.gemm",+ ):+ try:+ mod = importlib.import_module(module_name)+ except Exception:+ continue+ candidates.extend(+ [+ (mod, "gemm_a16wfp4_preshuffle", True),+ (mod, "gemm_a16wfp4_preshuffle_", False),+ ]+ )++ for holder, name, accepts_dict in candidates:+ fn = getattr(holder, name, None)+ if callable(fn):+ _DIRECT_HELPER = fn+ _DIRECT_HELPER_ACCEPTS_DICT = accepts_dict+ return+++ def _pick_shape_entry(m: int, n: int, k: int) -> dict[str, int | dict[str, object]]:+ shape = (m, n, k)+ cached = _SHAPE_CACHE.get(shape)+ if cached is not None:+ return cached++ tiles_bm16_n128 = _ceil_div(m, 16) * _ceil_div(n, 128)+ if m <= 32 or (m <= 128 and tiles_bm16_n128 < _LOW_UTIL_THRESHOLD):+ block_m = 8+ else:+ block_m = 16++ tiles_for_split = _ceil_div(m, block_m) * _ceil_div(n, 128)+ if m <= 32:+ if k >= 4096:+ ksplit = 7+ elif k >= 2048:+ ksplit = 4+ elif k >= 1536:+ ksplit = 3+ else:+ ksplit = 1+ elif k >= 2048 and tiles_for_split > _CU and tiles_for_split <= (_CU * 3) // 2:+ ksplit = 2+ elif k >= 7168 and (_CU // 2) <= tiles_for_split <= _CU:+ ksplit = 2+ elif block_m == 8 and k >= 2048 and (_CU // 2) <= tiles_for_split <= _CU:+ ksplit = 2+ else:+ ksplit = 1++ block_k = 256 if k <= (ksplit * 512) else 512+ block_n = 64 if (tiles_for_split * ksplit) < _LOW_UTIL_THRESHOLD else 128+ wgs = _ceil_div(m, block_m) * _ceil_div(n, block_n) * ksplit+ waves_per_eu = 2 if wgs > _CU else 1++ cfg = {+ "BLOCK_SIZE_M": block_m,+ "BLOCK_SIZE_N": block_n,+ "BLOCK_SIZE_K": block_k,+ "GROUP_SIZE_M": 1,+ "NUM_KSPLIT": ksplit,+ "SPLITK_BLOCK_SIZE": max(k // max(ksplit, 1), 64),+ "num_stages": 2,+ "num_warps": 4,+ "waves_per_eu": waves_per_eu,+ "matrix_instr_nonkdim": 16,+ "cache_modifier": ".cg",+ }++ entry = {+ "block_m": block_m,+ "block_n": block_n,+ "block_k": block_k,+ "ksplit": ksplit,+ "waves_per_eu": waves_per_eu,+ "cfg": cfg,+ }+ _SHAPE_CACHE[shape] = entry+ return entry+++ def _run_direct_session_path(+ a_bf16: torch.Tensor,+ b_shuffle: torch.Tensor,+ b_scale_sh: torch.Tensor,+ m: int,+ n: int,+ k: int,) -> torch.Tensor:- top = _run_quant_gemm(aiter, quant_func, dtypes, A[:SPLIT_CHUNK_M], B_shuffle, B_scale_sh)- bottom = _run_quant_gemm(aiter, quant_func, dtypes, A[SPLIT_CHUNK_M:], B_shuffle, B_scale_sh)- return torch.cat((top, bottom), dim=0)+ _resolve_direct_helper()+ if _DIRECT_HELPER is None:+ raise RuntimeError("a16 preshuffle helper not available")+ shape = (m, n, k)+ if not _DIRECT_SHAPE_SUPPORT.get(shape, True):+ raise RuntimeError(f"direct helper disabled for {shape}")- def custom_kernel(data: input_t) -> output_t:- """- Reference: MXFP4 per-1x32 quant on A; B_shuffle, B_scale_sh from generate_input.- gemm_a4w4 with bpreshuffle=True.- """- import aiter- from aiter import QuantType, dtypes+ entry = _pick_shape_entry(m, n, k)+ b_ps, s_ps = _get_cached_preshuffle_views(b_shuffle, b_scale_sh, n, k)+ cfg = entry["cfg"]+ try:+ config_arg = cfg+ if not _DIRECT_HELPER_ACCEPTS_DICT and _SERIALIZE_DICT is not None:+ config_arg = _SERIALIZE_DICT(_config_to_dict(cfg))+ return _DIRECT_HELPER(+ a_bf16,+ b_ps,+ s_ps,+ prequant=True,+ dtype=torch.bfloat16,+ y=None,+ config=config_arg,+ skip_reduce=False,+ )+ except Exception:+ _DIRECT_SHAPE_SUPPORT[shape] = False+ raise+++ def _run_fallback_gemm(+ a_q: torch.Tensor,+ b_shuffle: torch.Tensor,+ a_scale_sh: torch.Tensor,+ b_scale_sh: torch.Tensor,+ m: int,+ n: int,+ k: int,+ ) -> torch.Tensor:+ shape = (m, n, k)+ kernel_config = ASM_KERNEL_CONFIGS.get(shape)+ if ASM_GEMM is not None and kernel_config is not None:+ kernel_name, log2_k_split = kernel_config+ if _ASM_SHAPE_SUPPORT.get(shape, True):+ out = torch.empty((m, n), dtype=torch.bfloat16, device=a_q.device)+ try:+ ASM_GEMM(+ a_q,+ b_shuffle,+ a_scale_sh,+ b_scale_sh,+ out,+ kernelName=kernel_name,+ bpreshuffle=True,+ log2_k_split=log2_k_split,+ )+ _ASM_SHAPE_SUPPORT[shape] = True+ return out+ except Exception:+ _ASM_SHAPE_SUPPORT[shape] = False++ return aiter.gemm_a4w4(+ a_q,+ b_shuffle,+ a_scale_sh,+ b_scale_sh,+ dtype=dtypes.bf16,+ bpreshuffle=True,+ )+++ @torch.inference_mode()+ def custom_kernel(data: input_t) -> output_t:A, _B, _B_q, B_shuffle, B_scale_sh = data- A = A.contiguous()+ if not A.is_contiguous():+ A = A.contiguous()+m, k = A.shape+ n = B_shuffle.shape[0]+ shape = (m, n, k)- quant_func = aiter.get_triton_quant(QuantType.per_1x32)- if VARIANT == "shape_aware":- if m == SPLIT_ROW_M and k == SPLIT_ROW_K:- return _run_split_row_two_pass(aiter, quant_func, dtypes, A, B_shuffle, B_scale_sh)+ if shape in _FORCE_ASM_SHAPES:+ A_q, A_scale_sh = _get_cached_a_quant(A)+ return _run_fallback_gemm(A_q, B_shuffle, A_scale_sh, B_scale_sh, m, n, k)- return _run_quant_gemm(aiter, quant_func, dtypes, A, B_shuffle, B_scale_sh)+ try:+ return _run_direct_session_path(A, B_shuffle, B_scale_sh, m, n, k)+ except Exception:+ A_q, A_scale_sh = _get_cached_a_quant(A)+ return _run_fallback_gemm(A_q, B_shuffle, A_scale_sh, B_scale_sh, m, n, k)
scrolls · 428 diff lines total
Best evidence level for this revision: reported
JSON