submission 687604
.jonnss · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 534 lines, June 9 Researcher Reciprocity License v1.0.
Submission_v244.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-687604?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:1d9f8fa64b9a04c08433b83d665ee666ebd817a7f0e6a1747e40222ee538a3ac
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
_GET_SPLITK = NoneKernel source
Submission_v244.py534 lines
"""
Combined direct-path follow-up on top of Submission_v242.
This keeps the winning `(16, 16)` reduce tile from v242 and also forces
`waves_per_eu=2` for `(16, 2112, 7168)`.
"""
import importlib
import sys
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 = 16
_A_QUANT_CACHE = {}
_PRESHUFFLE_CACHE = {}
_SHAPE_CACHE = {}
_OUT_CACHE = {}
_PARTIAL_CACHE = {}
_DIRECT_INIT_DONE = False
_DIRECT_HELPER = None
_DIRECT_HELPER_ACCEPTS_DICT = True
_SERIALIZE_DICT = None
_DIRECT_KERNEL = None
_REDUCE_KERNEL = None
_GET_SPLITK = None
_TRITON = None
_DIRECT_KERNEL_SHAPE_SUPPORT = {}
_DIRECT_HELPER_SHAPE_SUPPORT = {}
_LOGGED_PATHS = {}
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 _trim_cache(cache: dict) -> None:
while len(cache) > _MAX_CACHE_ENTRIES:
cache.pop(next(iter(cache)))
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:
a_ref, a_ptr, a_version, a_q, a_scale_sh = cached
if a_ref() is a and a_ptr == a.data_ptr() and a_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)
stale_keys = [cache_key for cache_key, entry in _A_QUANT_CACHE.items() if entry[0]() is None]
for stale_key in stale_keys:
_A_QUANT_CACHE.pop(stale_key, None)
_trim_cache(_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, 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_u8, s_ps_u8 = 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_u8, s_ps_u8
b_ps_u8 = _view_dtype(b_shuffle, torch.uint8).contiguous().view(n // 16, k * 8).contiguous()
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,
)
stale_keys = [
cache_key
for cache_key, entry in _PRESHUFFLE_CACHE.items()
if entry[0]() is None or entry[1]() is None
]
for stale_key in stale_keys:
_PRESHUFFLE_CACHE.pop(stale_key, None)
_trim_cache(_PRESHUFFLE_CACHE)
return b_ps_u8, s_ps_u8
def _get_cached_output(device: torch.device, m: int, n: int) -> torch.Tensor:
key = (device.index or 0, m, n)
out = _OUT_CACHE.get(key)
if out is None or out.device != device or out.shape != (m, n):
out = torch.empty((m, n), dtype=torch.bfloat16, device=device)
_OUT_CACHE[key] = out
_trim_cache(_OUT_CACHE)
return out
def _get_cached_partials(device: torch.device, num_ksplit: int, m: int, n: int) -> torch.Tensor:
key = (device.index or 0, num_ksplit, m, n)
partials = _PARTIAL_CACHE.get(key)
if partials is None or partials.device != device or partials.shape != (num_ksplit, m, n):
partials = torch.empty((num_ksplit, m, n), dtype=torch.float32, device=device)
_PARTIAL_CACHE[key] = partials
_trim_cache(_PARTIAL_CACHE)
return partials
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_runtime() -> None:
global _DIRECT_INIT_DONE, _DIRECT_HELPER, _DIRECT_HELPER_ACCEPTS_DICT, _SERIALIZE_DICT
global _DIRECT_KERNEL, _REDUCE_KERNEL, _GET_SPLITK, _TRITON
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
try:
_TRITON = importlib.import_module("triton")
except Exception:
_TRITON = None
try:
kernel_mod = importlib.import_module("aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4")
_DIRECT_KERNEL = getattr(kernel_mod, "_gemm_a16wfp4_preshuffle_kernel", None)
except Exception:
_DIRECT_KERNEL = None
try:
reduce_mod = importlib.import_module("aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4")
_REDUCE_KERNEL = getattr(reduce_mod, "_gemm_afp4wfp4_reduce_kernel", None)
except Exception:
_REDUCE_KERNEL = None
try:
splitk_mod = importlib.import_module("aiter.ops.triton.gemm.basic.gemm_afp4wfp4")
_GET_SPLITK = getattr(splitk_mod, "get_splitk", None)
except Exception:
_GET_SPLITK = None
candidates = []
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_", False),
(mod, "gemm_a16wfp4_preshuffle", True),
]
)
candidates.extend(
[
(aiter, "gemm_a16wfp4_preshuffle_", False),
(aiter, "gemm_a16wfp4_preshuffle", True),
]
)
for holder, name, accepts_dict in candidates:
fn = getattr(holder, name, None)
if callable(fn):
_DIRECT_HELPER = fn
_DIRECT_HELPER_ACCEPTS_DICT = accepts_dict
break
def _pick_shape_entry(m: int, n: int, k: 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
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": 2 if wgs > _CU else 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
}
if shape == (16, 2112, 7168):
cfg["waves_per_eu"] = 2
if shape == (64, 7168, 2048):
cfg["waves_per_eu"] = 1
entry = {"cfg": cfg}
_SHAPE_CACHE[shape] = entry
_trim_cache(_SHAPE_CACHE)
return entry
def _prepare_helper_cfg(m: int, n: int, k: int) -> dict[str, object]:
cfg = dict(_config_to_dict(_pick_shape_entry(m, n, k)["cfg"]))
if cfg["NUM_KSPLIT"] > 1 and _GET_SPLITK is not None:
splitk_block_size, block_size_k, num_ksplit = _GET_SPLITK(
k, cfg["BLOCK_SIZE_K"], cfg["NUM_KSPLIT"]
)
cfg["SPLITK_BLOCK_SIZE"] = splitk_block_size
cfg["BLOCK_SIZE_K"] = block_size_k
cfg["NUM_KSPLIT"] = num_ksplit
if _TRITON is not None and cfg["BLOCK_SIZE_K"] >= 2 * k:
cfg["BLOCK_SIZE_K"] = int(_TRITON.next_power_of_2(2 * k))
cfg["SPLITK_BLOCK_SIZE"] = 2 * k
cfg["NUM_KSPLIT"] = 1
cfg["BLOCK_SIZE_N"] = max(cfg["BLOCK_SIZE_N"], 32)
if cfg["NUM_KSPLIT"] <= 1:
cfg["NUM_KSPLIT"] = 1
cfg["SPLITK_BLOCK_SIZE"] = 2 * k
return cfg
def _prepare_direct_cfg(m: int, n: int, k: int, runtime_k: int) -> dict[str, object]:
cfg = dict(_config_to_dict(_pick_shape_entry(m, n, k)["cfg"]))
if cfg["NUM_KSPLIT"] > 1 and _GET_SPLITK is not None:
splitk_block_size, block_size_k, num_ksplit = _GET_SPLITK(
runtime_k, cfg["BLOCK_SIZE_K"], cfg["NUM_KSPLIT"]
)
cfg["SPLITK_BLOCK_SIZE"] = splitk_block_size
cfg["BLOCK_SIZE_K"] = block_size_k
cfg["NUM_KSPLIT"] = num_ksplit
if _TRITON is not None and cfg["BLOCK_SIZE_K"] >= 2 * runtime_k:
cfg["BLOCK_SIZE_K"] = int(_TRITON.next_power_of_2(2 * runtime_k))
cfg["SPLITK_BLOCK_SIZE"] = 2 * runtime_k
cfg["NUM_KSPLIT"] = 1
cfg["BLOCK_SIZE_N"] = max(cfg["BLOCK_SIZE_N"], 32)
if cfg["NUM_KSPLIT"] <= 1:
cfg["NUM_KSPLIT"] = 1
cfg["SPLITK_BLOCK_SIZE"] = 2 * runtime_k
return cfg
def _cfg_brief(cfg: dict[str, object]) -> str:
return (
f"bm={cfg['BLOCK_SIZE_M']},bn={cfg['BLOCK_SIZE_N']},bk={cfg['BLOCK_SIZE_K']},"
f"sp={cfg['NUM_KSPLIT']},sb={cfg['SPLITK_BLOCK_SIZE']},st={cfg['num_stages']},"
f"wp={cfg['num_warps']},wpe={cfg['waves_per_eu']}"
)
def _emit_path(shape: tuple[int, int, int], path: str, cfg: dict[str, object], detail: str = "") -> None:
previous = _LOGGED_PATHS.get(shape)
if previous is not None:
return
_LOGGED_PATHS[shape] = path
suffix = f" {detail}" if detail else ""
print(
f"[amd2-mm-v244] shape={shape} path={path} {_cfg_brief(cfg)}{suffix}",
file=sys.stderr,
flush=True,
)
def _run_direct_kernel_path(
a_bf16: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
m: int,
n: int,
k: int,
) -> torch.Tensor:
_resolve_runtime()
shape = (m, n, k)
if not _DIRECT_KERNEL_SHAPE_SUPPORT.get(shape, True):
raise RuntimeError(f"direct kernel disabled for {shape}")
if _DIRECT_KERNEL is None or _TRITON is None:
raise RuntimeError("direct Triton preshuffle kernel unavailable")
b_ps_u8, s_ps_u8 = _get_cached_preshuffle_views(b_shuffle, b_scale_sh, n, k)
runtime_n = b_ps_u8.shape[0] * 16
runtime_k = b_ps_u8.shape[1] // 16
cfg = _prepare_direct_cfg(m, n, k, runtime_k)
if cfg["NUM_KSPLIT"] > 1 and _REDUCE_KERNEL is None:
raise RuntimeError("direct Triton reduce kernel unavailable")
y = _get_cached_output(a_bf16.device, m, runtime_n)
if cfg["NUM_KSPLIT"] > 1:
y_pp = _get_cached_partials(a_bf16.device, int(cfg["NUM_KSPLIT"]), m, runtime_n)
out = y_pp
else:
y_pp = None
out = y
grid = lambda meta: ( # noqa: E731
(
meta["NUM_KSPLIT"]
* _ceil_div(m, int(meta["BLOCK_SIZE_M"]))
* _ceil_div(runtime_n, int(meta["BLOCK_SIZE_N"]))
),
)
try:
_DIRECT_KERNEL[grid](
a_bf16,
b_ps_u8,
out,
s_ps_u8,
m,
runtime_n,
runtime_k,
a_bf16.stride(0),
a_bf16.stride(1),
b_ps_u8.stride(0),
b_ps_u8.stride(1),
0 if y_pp is None else y_pp.stride(0),
y.stride(0) if y_pp is None else y_pp.stride(1),
y.stride(1) if y_pp is None else y_pp.stride(2),
s_ps_u8.stride(0),
s_ps_u8.stride(1),
PREQUANT=True,
**cfg,
)
if y_pp is not None:
actual_ksplit = int(_TRITON.cdiv(runtime_k, int(cfg["SPLITK_BLOCK_SIZE"]) // 2))
grid_reduce = (_ceil_div(m, 16), _ceil_div(runtime_n, 16))
_REDUCE_KERNEL[grid_reduce](
y_pp,
y,
m,
runtime_n,
y_pp.stride(0),
y_pp.stride(1),
y_pp.stride(2),
y.stride(0),
y.stride(1),
16,
16,
actual_ksplit,
int(_TRITON.next_power_of_2(int(cfg["NUM_KSPLIT"]))),
)
return y
except Exception:
_DIRECT_KERNEL_SHAPE_SUPPORT[shape] = False
raise
def _run_direct_helper_path(
a_bf16: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
m: int,
n: int,
k: int,
) -> torch.Tensor:
_resolve_runtime()
if _DIRECT_HELPER is None:
raise RuntimeError("direct helper unavailable")
shape = (m, n, k)
if not _DIRECT_HELPER_SHAPE_SUPPORT.get(shape, True):
raise RuntimeError(f"direct helper disabled for {shape}")
cfg = _prepare_helper_cfg(m, n, k)
b_ps_u8, s_ps_u8 = _get_cached_preshuffle_views(b_shuffle, b_scale_sh, n, k)
y = _get_cached_output(a_bf16.device, m, n)
try:
config_arg = cfg
if not _DIRECT_HELPER_ACCEPTS_DICT and _SERIALIZE_DICT is not None:
config_arg = _SERIALIZE_DICT(cfg)
return _DIRECT_HELPER(
a_bf16,
b_ps_u8,
s_ps_u8,
prequant=True,
dtype=torch.bfloat16,
y=y,
config=config_arg,
skip_reduce=False,
)
except Exception:
_DIRECT_HELPER_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:
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)
helper_cfg = _prepare_helper_cfg(m, n, k)
b_ps_u8, _s_ps_u8 = _get_cached_preshuffle_views(B_shuffle, B_scale_sh, n, k)
runtime_n = b_ps_u8.shape[0] * 16
runtime_k = b_ps_u8.shape[1] // 16
direct_cfg = _prepare_direct_cfg(m, n, k, runtime_k)
try:
out = _run_direct_kernel_path(A, B_shuffle, B_scale_sh, m, n, k)
_emit_path(shape, "direct", direct_cfg, f"rn={runtime_n},rk={runtime_k}")
return out
except Exception as direct_exc:
direct_detail = repr(direct_exc)
try:
out = _run_direct_helper_path(A, B_shuffle, B_scale_sh, m, n, k)
_emit_path(shape, "helper", helper_cfg, f"rn={runtime_n},rk={runtime_k},direct={direct_detail}")
return out
except Exception as helper_exc:
A_q, A_scale_sh = _get_cached_a_quant(A)
_emit_path(
shape,
"fallback",
helper_cfg,
f"rn={runtime_n},rk={runtime_k},direct={direct_detail},helper={repr(helper_exc)}",
)
return _run_fallback_gemm(A_q, B_shuffle, A_scale_sh, B_scale_sh, m, n, k)
scrolls · 534 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 679226.
"""- SESSION.md-guided A16 preshuffle rebuild.+ Combined direct-path follow-up on top of Submission_v242.- 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.+ This keeps the winning `(16, 16)` reduce tile from v242 and also forces+ `waves_per_eu=2` for `(16, 2112, 7168)`."""import importlib+ import sysimport weakrefimport aiter⋯ 7 unchanged lines_CU = 256_LOW_UTIL_THRESHOLD = (_CU * 3) // 4+ _MAX_CACHE_ENTRIES = 16- _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,- ],- ] = {}+ _A_QUANT_CACHE = {}+ _PRESHUFFLE_CACHE = {}+ _SHAPE_CACHE = {}+ _OUT_CACHE = {}+ _PARTIAL_CACHE = {}_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]]] = {}+ _DIRECT_KERNEL = None+ _REDUCE_KERNEL = None+ _GET_SPLITK = None+ _TRITON = None+ _DIRECT_KERNEL_SHAPE_SUPPORT = {}+ _DIRECT_HELPER_SHAPE_SUPPORT = {}+ _LOGGED_PATHS = {}- 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): ("f4gemm_bf16_per1x32Fp4_BpreShuffle_192x128", None),- }- _ASM_SHAPE_SUPPORT: dict[tuple[int, int, int], bool] = {}-def _ceil_div(a: int, b: int) -> int:return (a + b - 1) // b⋯ 4 unchanged linesreturn tensor.view(dtype)+ def _trim_cache(cache: dict) -> None:+ while len(cache) > _MAX_CACHE_ENTRIES:+ cache.pop(next(iter(cache)))++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)⋯ 4 unchanged lineskey = (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:+ a_ref, a_ptr, a_version, a_q, a_scale_sh = cached+ if a_ref() is a and a_ptr == a.data_ptr() and a_version == a._version:return a_q, a_scale_sha_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)))-+ stale_keys = [cache_key for cache_key, entry in _A_QUANT_CACHE.items() if entry[0]() is None]+ for stale_key in stale_keys:+ _A_QUANT_CACHE.pop(stale_key, None)+ _trim_cache(_A_QUANT_CACHE)return a_q, a_scale_sh⋯ 3 unchanged linesn: int,k: int,) -> tuple[torch.Tensor, torch.Tensor]:- key = (b_shuffle.data_ptr(), b_scale_sh.data_ptr(), n * 1_000_000 + k)+ key = (b_shuffle.data_ptr(), b_scale_sh.data_ptr(), n, 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+ b_ref, s_ref, b_ptr, s_ptr, b_version, s_version, b_ps_u8, s_ps_u8 = cachedif (b_ref() is b_shuffleand s_ref() is b_scale_sh⋯ 2 unchanged linesand b_version == b_shuffle._versionand s_version == b_scale_sh._version):- return b_ps, s_ps+ return b_ps_u8, s_ps_u8b_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()⋯ 7 unchanged linesb_ps_u8,s_ps_u8,)+ stale_keys = [+ cache_key+ for cache_key, entry in _PRESHUFFLE_CACHE.items()+ if entry[0]() is None or entry[1]() is None+ ]+ for stale_key in stale_keys:+ _PRESHUFFLE_CACHE.pop(stale_key, None)+ _trim_cache(_PRESHUFFLE_CACHE)+ return 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 _get_cached_output(device: torch.device, m: int, n: int) -> torch.Tensor:+ key = (device.index or 0, m, n)+ out = _OUT_CACHE.get(key)+ if out is None or out.device != device or out.shape != (m, n):+ out = torch.empty((m, n), dtype=torch.bfloat16, device=device)+ _OUT_CACHE[key] = out+ _trim_cache(_OUT_CACHE)+ return out+ def _get_cached_partials(device: torch.device, num_ksplit: int, m: int, n: int) -> torch.Tensor:+ key = (device.index or 0, num_ksplit, m, n)+ partials = _PARTIAL_CACHE.get(key)+ if partials is None or partials.device != device or partials.shape != (num_ksplit, m, n):+ partials = torch.empty((num_ksplit, m, n), dtype=torch.float32, device=device)+ _PARTIAL_CACHE[key] = partials+ _trim_cache(_PARTIAL_CACHE)+ return partials++def _config_to_dict(base) -> dict[str, object]:if base is None:return {}⋯ 13 unchanged linesreturn {}- def _resolve_direct_helper() -> None:+ def _resolve_runtime() -> None:global _DIRECT_INIT_DONE, _DIRECT_HELPER, _DIRECT_HELPER_ACCEPTS_DICT, _SERIALIZE_DICT+ global _DIRECT_KERNEL, _REDUCE_KERNEL, _GET_SPLITK, _TRITONif _DIRECT_INIT_DONE:return_DIRECT_INIT_DONE = True⋯ 4 unchanged linesexcept Exception:_SERIALIZE_DICT = None- candidates: list[tuple[object, str, bool]] = [- (aiter, "gemm_a16wfp4_preshuffle", True),- (aiter, "gemm_a16wfp4_preshuffle_", False),- ]+ try:+ _TRITON = importlib.import_module("triton")+ except Exception:+ _TRITON = None++ try:+ kernel_mod = importlib.import_module("aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4")+ _DIRECT_KERNEL = getattr(kernel_mod, "_gemm_a16wfp4_preshuffle_kernel", None)+ except Exception:+ _DIRECT_KERNEL = None++ try:+ reduce_mod = importlib.import_module("aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4")+ _REDUCE_KERNEL = getattr(reduce_mod, "_gemm_afp4wfp4_reduce_kernel", None)+ except Exception:+ _REDUCE_KERNEL = None++ try:+ splitk_mod = importlib.import_module("aiter.ops.triton.gemm.basic.gemm_afp4wfp4")+ _GET_SPLITK = getattr(splitk_mod, "get_splitk", None)+ except Exception:+ _GET_SPLITK = None++ candidates = []for module_name in ("aiter.ops.triton.gemm.basic.gemm_a16wfp4","aiter.ops.triton.gemm.gemm_a16wfp4",⋯ 6 unchanged linescontinuecandidates.extend([- (mod, "gemm_a16wfp4_preshuffle", True),(mod, "gemm_a16wfp4_preshuffle_", False),+ (mod, "gemm_a16wfp4_preshuffle", True),])+ candidates.extend(+ [+ (aiter, "gemm_a16wfp4_preshuffle_", False),+ (aiter, "gemm_a16wfp4_preshuffle", True),+ ]+ )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+ break- def _pick_shape_entry(m: int, n: int, k: int) -> dict[str, int | dict[str, object]]:+ def _pick_shape_entry(m: int, n: int, k: int) -> dict[str, object]:shape = (m, n, k)cached = _SHAPE_CACHE.get(shape)if cached is not None:⋯ 27 unchanged linesblock_k = 256 if k <= (ksplit * 512) else 512block_n = 64 if (tiles_for_split * ksplit) < _LOW_UTIL_THRESHOLD else 128wgs = _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,⋯ 3 unchanged lines"SPLITK_BLOCK_SIZE": max(k // max(ksplit, 1), 64),"num_stages": 2,"num_warps": 4,- "waves_per_eu": waves_per_eu,+ "waves_per_eu": 2 if wgs > _CU else 1,"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,- }+ if shape == (16, 2112, 7168):+ cfg["waves_per_eu"] = 2+ if shape == (64, 7168, 2048):+ cfg["waves_per_eu"] = 1+ entry = {"cfg": cfg}_SHAPE_CACHE[shape] = entry+ _trim_cache(_SHAPE_CACHE)return entry- def _run_direct_session_path(+ def _prepare_helper_cfg(m: int, n: int, k: int) -> dict[str, object]:+ cfg = dict(_config_to_dict(_pick_shape_entry(m, n, k)["cfg"]))+ if cfg["NUM_KSPLIT"] > 1 and _GET_SPLITK is not None:+ splitk_block_size, block_size_k, num_ksplit = _GET_SPLITK(+ k, cfg["BLOCK_SIZE_K"], cfg["NUM_KSPLIT"]+ )+ cfg["SPLITK_BLOCK_SIZE"] = splitk_block_size+ cfg["BLOCK_SIZE_K"] = block_size_k+ cfg["NUM_KSPLIT"] = num_ksplit++ if _TRITON is not None and cfg["BLOCK_SIZE_K"] >= 2 * k:+ cfg["BLOCK_SIZE_K"] = int(_TRITON.next_power_of_2(2 * k))+ cfg["SPLITK_BLOCK_SIZE"] = 2 * k+ cfg["NUM_KSPLIT"] = 1++ cfg["BLOCK_SIZE_N"] = max(cfg["BLOCK_SIZE_N"], 32)+ if cfg["NUM_KSPLIT"] <= 1:+ cfg["NUM_KSPLIT"] = 1+ cfg["SPLITK_BLOCK_SIZE"] = 2 * k+ return cfg+++ def _prepare_direct_cfg(m: int, n: int, k: int, runtime_k: int) -> dict[str, object]:+ cfg = dict(_config_to_dict(_pick_shape_entry(m, n, k)["cfg"]))+ if cfg["NUM_KSPLIT"] > 1 and _GET_SPLITK is not None:+ splitk_block_size, block_size_k, num_ksplit = _GET_SPLITK(+ runtime_k, cfg["BLOCK_SIZE_K"], cfg["NUM_KSPLIT"]+ )+ cfg["SPLITK_BLOCK_SIZE"] = splitk_block_size+ cfg["BLOCK_SIZE_K"] = block_size_k+ cfg["NUM_KSPLIT"] = num_ksplit++ if _TRITON is not None and cfg["BLOCK_SIZE_K"] >= 2 * runtime_k:+ cfg["BLOCK_SIZE_K"] = int(_TRITON.next_power_of_2(2 * runtime_k))+ cfg["SPLITK_BLOCK_SIZE"] = 2 * runtime_k+ cfg["NUM_KSPLIT"] = 1++ cfg["BLOCK_SIZE_N"] = max(cfg["BLOCK_SIZE_N"], 32)+ if cfg["NUM_KSPLIT"] <= 1:+ cfg["NUM_KSPLIT"] = 1+ cfg["SPLITK_BLOCK_SIZE"] = 2 * runtime_k+ return cfg+++ def _cfg_brief(cfg: dict[str, object]) -> str:+ return (+ f"bm={cfg['BLOCK_SIZE_M']},bn={cfg['BLOCK_SIZE_N']},bk={cfg['BLOCK_SIZE_K']},"+ f"sp={cfg['NUM_KSPLIT']},sb={cfg['SPLITK_BLOCK_SIZE']},st={cfg['num_stages']},"+ f"wp={cfg['num_warps']},wpe={cfg['waves_per_eu']}"+ )+++ def _emit_path(shape: tuple[int, int, int], path: str, cfg: dict[str, object], detail: str = "") -> None:+ previous = _LOGGED_PATHS.get(shape)+ if previous is not None:+ return+ _LOGGED_PATHS[shape] = path+ suffix = f" {detail}" if detail else ""+ print(+ f"[amd2-mm-v244] shape={shape} path={path} {_cfg_brief(cfg)}{suffix}",+ file=sys.stderr,+ flush=True,+ )+++ def _run_direct_kernel_path(a_bf16: torch.Tensor,b_shuffle: torch.Tensor,b_scale_sh: torch.Tensor,⋯ 1 unchanged linesn: int,k: int,) -> torch.Tensor:- _resolve_direct_helper()+ _resolve_runtime()++ shape = (m, n, k)+ if not _DIRECT_KERNEL_SHAPE_SUPPORT.get(shape, True):+ raise RuntimeError(f"direct kernel disabled for {shape}")++ if _DIRECT_KERNEL is None or _TRITON is None:+ raise RuntimeError("direct Triton preshuffle kernel unavailable")++ b_ps_u8, s_ps_u8 = _get_cached_preshuffle_views(b_shuffle, b_scale_sh, n, k)+ runtime_n = b_ps_u8.shape[0] * 16+ runtime_k = b_ps_u8.shape[1] // 16+ cfg = _prepare_direct_cfg(m, n, k, runtime_k)+ if cfg["NUM_KSPLIT"] > 1 and _REDUCE_KERNEL is None:+ raise RuntimeError("direct Triton reduce kernel unavailable")+ y = _get_cached_output(a_bf16.device, m, runtime_n)++ if cfg["NUM_KSPLIT"] > 1:+ y_pp = _get_cached_partials(a_bf16.device, int(cfg["NUM_KSPLIT"]), m, runtime_n)+ out = y_pp+ else:+ y_pp = None+ out = y++ grid = lambda meta: ( # noqa: E731+ (+ meta["NUM_KSPLIT"]+ * _ceil_div(m, int(meta["BLOCK_SIZE_M"]))+ * _ceil_div(runtime_n, int(meta["BLOCK_SIZE_N"]))+ ),+ )++ try:+ _DIRECT_KERNEL[grid](+ a_bf16,+ b_ps_u8,+ out,+ s_ps_u8,+ m,+ runtime_n,+ runtime_k,+ a_bf16.stride(0),+ a_bf16.stride(1),+ b_ps_u8.stride(0),+ b_ps_u8.stride(1),+ 0 if y_pp is None else y_pp.stride(0),+ y.stride(0) if y_pp is None else y_pp.stride(1),+ y.stride(1) if y_pp is None else y_pp.stride(2),+ s_ps_u8.stride(0),+ s_ps_u8.stride(1),+ PREQUANT=True,+ **cfg,+ )++ if y_pp is not None:+ actual_ksplit = int(_TRITON.cdiv(runtime_k, int(cfg["SPLITK_BLOCK_SIZE"]) // 2))+ grid_reduce = (_ceil_div(m, 16), _ceil_div(runtime_n, 16))+ _REDUCE_KERNEL[grid_reduce](+ y_pp,+ y,+ m,+ runtime_n,+ y_pp.stride(0),+ y_pp.stride(1),+ y_pp.stride(2),+ y.stride(0),+ y.stride(1),+ 16,+ 16,+ actual_ksplit,+ int(_TRITON.next_power_of_2(int(cfg["NUM_KSPLIT"]))),+ )+ return y+ except Exception:+ _DIRECT_KERNEL_SHAPE_SUPPORT[shape] = False+ raise+++ def _run_direct_helper_path(+ a_bf16: torch.Tensor,+ b_shuffle: torch.Tensor,+ b_scale_sh: torch.Tensor,+ m: int,+ n: int,+ k: int,+ ) -> torch.Tensor:+ _resolve_runtime()if _DIRECT_HELPER is None:- raise RuntimeError("a16 preshuffle helper not available")+ raise RuntimeError("direct helper unavailable")shape = (m, n, k)- if not _DIRECT_SHAPE_SUPPORT.get(shape, True):+ if not _DIRECT_HELPER_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"]+ cfg = _prepare_helper_cfg(m, n, k)+ b_ps_u8, s_ps_u8 = _get_cached_preshuffle_views(b_shuffle, b_scale_sh, n, k)+ y = _get_cached_output(a_bf16.device, m, n)try:config_arg = cfgif not _DIRECT_HELPER_ACCEPTS_DICT and _SERIALIZE_DICT is not None:- config_arg = _SERIALIZE_DICT(_config_to_dict(cfg))+ config_arg = _SERIALIZE_DICT(cfg)return _DIRECT_HELPER(a_bf16,- b_ps,- s_ps,+ b_ps_u8,+ s_ps_u8,prequant=True,dtype=torch.bfloat16,- y=None,+ y=y,config=config_arg,skip_reduce=False,)except Exception:- _DIRECT_SHAPE_SUPPORT[shape] = False+ _DIRECT_HELPER_SHAPE_SUPPORT[shape] = Falseraise⋯ 6 unchanged linesn: 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,⋯ 12 unchanged linesm, k = A.shapen = B_shuffle.shape[0]+ shape = (m, n, k)+ helper_cfg = _prepare_helper_cfg(m, n, k)+ b_ps_u8, _s_ps_u8 = _get_cached_preshuffle_views(B_shuffle, B_scale_sh, n, k)+ runtime_n = b_ps_u8.shape[0] * 16+ runtime_k = b_ps_u8.shape[1] // 16+ direct_cfg = _prepare_direct_cfg(m, n, k, runtime_k)try:- return _run_direct_session_path(A, B_shuffle, B_scale_sh, m, n, k)- except Exception:+ out = _run_direct_kernel_path(A, B_shuffle, B_scale_sh, m, n, k)+ _emit_path(shape, "direct", direct_cfg, f"rn={runtime_n},rk={runtime_k}")+ return out+ except Exception as direct_exc:+ direct_detail = repr(direct_exc)++ try:+ out = _run_direct_helper_path(A, B_shuffle, B_scale_sh, m, n, k)+ _emit_path(shape, "helper", helper_cfg, f"rn={runtime_n},rk={runtime_k},direct={direct_detail}")+ return out+ except Exception as helper_exc:A_q, A_scale_sh = _get_cached_a_quant(A)+ _emit_path(+ shape,+ "fallback",+ helper_cfg,+ f"rn={runtime_n},rk={runtime_k},direct={direct_detail},helper={repr(helper_exc)}",+ )return _run_fallback_gemm(A_q, B_shuffle, A_scale_sh, B_scale_sh, m, n, k)
scrolls · 580 diff lines total
Best evidence level for this revision: reported
JSON