submission 604221
Roshan Rateria · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 726 lines, June 9 Researcher Reciprocity License v1.0.
submission_best_till_now.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-604221?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:6d14f46766f4c0db8f8d31a301733c8ac749883b4b35ce6ce156286e338e9bc1
license declaredunknown
license concludedunknown
authorsRoshan Rateria
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
autotune
def _autotune_kernel_splitk(A_q, B_shuffle, A_scale_sh, B_scale_sh, out_buf, kernelName, splitK):fp4
Optimized MXFP4 GEMM for AMD MI355X.num-warps = 1
NUM_WARPS = 1split-k
- It calls get_GEMM_config to look up kernelName + splitK from CSVstages = 1
NUM_STAGES = 1tile-m = 64
BLOCK_SIZE_M = 64tile-n = 32
BLOCK_SIZE_N = 32Kernel source
submission_best_till_now.py726 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
Optimized MXFP4 GEMM for AMD MI355X.
Key insight from aiter source (gemm_op_a4w4.py):
- gemm_a4w4 allocates out = torch.empty(((m+31)//32*32, n), ...) every call
- It calls get_GEMM_config to look up kernelName + splitK from CSV
- gemm_a4w4_asm takes the ORIGINAL (unpadded) A and A_scale -- kernel handles padding
- For our benchmark shapes, no tuned config exists -> kernelName="", splitK=0
Optimization: pre-allocate output buffer per shape, reuse across calls.
Use get_GEMM_config to get correct kernelName/splitK (avoids CSV re-read).
"""
from task import input_t, output_t
import atexit
import json
import os
import weakref
_INLINE_ARCH = os.environ.get("MXFP4_INLINE_ARCH", "gfx950")
_USE_INLINE = os.environ.get("MXFP4_USE_INLINE", "0") != "0"
if _USE_INLINE:
os.environ.setdefault("PYTORCH_ROCM_ARCH", _INLINE_ARCH)
os.environ.setdefault("CXX", "clang++")
import torch
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
from aiter.ops.gemm_op_a4w4 import (
gemm_a4w4_asm,
gemm_a4w4_blockscale,
get_GEMM_config,
)
_FP4X2 = dtypes.fp4x2
_FP8E8M0 = dtypes.fp8_e8m0
_BF16 = dtypes.bf16
_quant = dynamic_mxfp4_quant
_shuffle = e8m0_shuffle
_gemm_asm = gemm_a4w4_asm
_gemm_blk = gemm_a4w4_blockscale
_get_cfg = get_GEMM_config
# Per-shape cache: (M, N, K) -> (padded_M, out_buf, use_asm, kernelName, splitK)
_cache: dict = {}
_warmed = False
_USE_GRAPH = os.environ.get("MXFP4_USE_GRAPH", "1") != "0"
_GRAPH_ENABLED = _USE_GRAPH and hasattr(torch.cuda, "CUDAGraph") and torch.cuda.is_available()
_GRAPH_CACHE: dict = {}
_GRAPH_BLACKLIST: set = set()
_GRAPH_HITS = 0
_GRAPH_MISSES = 0
_GRAPH_FAILS = 0
_DEBUG = os.environ.get("MXFP4_DEBUG", "0") != "0"
_FORCE_KERNEL = os.environ.get("MXFP4_FORCE_KERNEL", "").strip()
_FORCE_SPLITK = os.environ.get("MXFP4_FORCE_SPLITK", "").strip()
_AUTO_KERNEL = os.environ.get("MXFP4_AUTO_KERNEL", "1") != "0"
_KERNEL_32 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_KERNEL_192 = "_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_192x128E"
_AUTOTUNE_KERNEL = os.environ.get("MXFP4_AUTOTUNE_KERNEL", "0") != "0"
_AUTOTUNE_SPLITK = os.environ.get("MXFP4_AUTOTUNE_SPLITK", "0") != "0"
_AUTOTUNE_ALWAYS = os.environ.get("MXFP4_AUTOTUNE_ALWAYS", "0") != "0"
_AUTOTUNE_REPS = max(1, int(os.environ.get("MXFP4_AUTOTUNE_REPS", "5")))
_KERNEL_CANDS_ENV = os.environ.get("MXFP4_KERNEL_CANDS", "32x128,192x128")
_SPLITK_CANDS_ENV = os.environ.get("MXFP4_SPLITK_CANDS", "0,1,2")
_PRINT_TUNED_MAP = os.environ.get("MXFP4_PRINT_TUNED_MAP", "1") != "0"
_TUNED_MAP_PRINTED: set = set()
_F4GEMM_KERNELS: list = []
_TUNED_MAP: dict = {
# Default tuned map from MI355X leaderboard run
(4, 2880, 512): (_KERNEL_192, 2),
(16, 2112, 7168): (_KERNEL_32, 2),
(32, 4096, 512): (_KERNEL_192, 2),
(32, 2880, 512): (_KERNEL_32, 2),
(64, 7168, 2048): (_KERNEL_32, 1),
(256, 3072, 1536): (_KERNEL_32, 2),
}
_TUNED_MAP_ENV = os.environ.get("MXFP4_TUNED_MAP", "").strip()
if _TUNED_MAP_ENV:
_TUNED_MAP.clear()
try:
# Expect JSON: {"M,N,K":{"kernel":"...","splitK":0}, ...}
raw = json.loads(_TUNED_MAP_ENV)
for k, v in raw.items():
m, n, kk = (int(x) for x in k.split(","))
_TUNED_MAP[(m, n, kk)] = (v.get("kernel", ""), int(v.get("splitK", 0)))
except Exception:
if _DEBUG:
print("[mxfp4] failed to parse MXFP4_TUNED_MAP, ignoring")
if _DEBUG:
def _print_graph_stats():
print(
f"[mxfp4] graph hits={_GRAPH_HITS} misses={_GRAPH_MISSES} fails={_GRAPH_FAILS}"
)
atexit.register(_print_graph_stats)
if _PRINT_TUNED_MAP and (_AUTOTUNE_KERNEL or _AUTOTUNE_SPLITK):
def _print_tuned_map():
if _TUNED_MAP:
print(f"[mxfp4] tuned map: {_TUNED_MAP}")
atexit.register(_print_tuned_map)
def _load_f4gemm_kernels():
global _F4GEMM_KERNELS
if _F4GEMM_KERNELS:
return _F4GEMM_KERNELS
try:
aiter_root = os.path.abspath(os.path.join(os.path.dirname(aiter.__file__), os.pardir))
csv_path = os.path.join(
aiter_root, "hsa", _INLINE_ARCH, "f4gemm", "f4gemm_bf16_per1x32Fp4.csv"
)
if not os.path.exists(csv_path):
return _F4GEMM_KERNELS
kernels = []
with open(csv_path, "r", encoding="utf-8") as f:
header = f.readline()
for line in f:
parts = line.strip().split(",")
if len(parts) < 5:
continue
# tile_M,tile_N,splitK,bpreshuffle,knl_name,...
bpreshuffle = parts[3].strip()
knl_name = parts[4].strip()
if bpreshuffle != "1":
continue
if "_BpreShuffle_" not in knl_name:
continue
kernels.append(knl_name)
# de-dup while preserving order
seen = set()
uniq = []
for k in kernels:
if k not in seen:
uniq.append(k)
seen.add(k)
_F4GEMM_KERNELS = uniq
return _F4GEMM_KERNELS
except Exception:
return _F4GEMM_KERNELS
_INLINE_EXEC = os.environ.get("MXFP4_INLINE_EXEC", "0") != "0"
_INLINE_AVAILABLE = False
_INLINE_WARNED = False
_inline_mod = None
_FUSED_SHUFFLE = os.environ.get("MXFP4_FUSED_SHUFFLE", "1") != "0"
_AQ_CACHE: dict = {}
_CACHE_AQ = os.environ.get("MXFP4_CACHE_AQ", "1") != "0"
_AQ_REUSE_MAX = max(1, int(os.environ.get("MXFP4_CACHE_AQ_MAX", "16")))
_AQ_REUSE_CACHE: dict = {}
_AQ_REUSE_ORDER: list = []
if _FUSED_SHUFFLE:
import triton
import triton.language as tl
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
@triton.heuristics(
{
"EVEN_M_N": lambda args: args["M"] % args["BLOCK_SIZE_M"] == 0
and args["N"] % (args["BLOCK_SIZE_N"] * args["NUM_ITER"]) == 0,
}
)
@triton.jit
def _dynamic_mxfp4_quant_kernel_shuffled(
x_ptr,
x_fp4_ptr,
bs_ptr,
stride_x_m_in,
stride_x_n_in,
stride_x_fp4_m_in,
stride_x_fp4_n_in,
stride_bs_m_in,
stride_bs_n_in,
M,
N,
scaleN,
scaleM_pad,
scaleN_pad,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
NUM_ITER: tl.constexpr,
NUM_STAGES: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
EVEN_M_N: tl.constexpr,
):
pid_m = tl.program_id(0)
start_n = tl.program_id(1) * NUM_ITER
stride_x_m = tl.cast(stride_x_m_in, tl.int64)
stride_x_n = tl.cast(stride_x_n_in, tl.int64)
stride_x_fp4_m = tl.cast(stride_x_fp4_m_in, tl.int64)
stride_x_fp4_n = tl.cast(stride_x_fp4_n_in, tl.int64)
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):
x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
if EVEN_M_N:
x = tl.load(x_ptr + x_offs, cache_modifier=".cg").to(tl.float32)
else:
x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
x = tl.load(x_ptr + x_offs, mask=x_mask, cache_modifier=".cg").to(
tl.float32
)
out_tensor, bs_e8m0 = _mxfp4_quant_op(
x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE
)
out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
out_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
out_offs = (
out_offs_m[:, None] * stride_x_fp4_m
+ out_offs_n[None, :] * stride_x_fp4_n
)
if EVEN_M_N:
tl.store(x_fp4_ptr + out_offs, out_tensor)
else:
out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)
bs_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
bs_offs_n = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
# Shuffle scale layout to match e8m0_shuffle
bs_offs_0 = bs_offs_m[:, None] // 32
bs_offs_1 = bs_offs_m[:, None] % 32
bs_offs_2 = bs_offs_1 % 16
bs_offs_1 = bs_offs_1 // 16
bs_offs_3 = bs_offs_n[None, :] // 8
bs_offs_4 = bs_offs_n[None, :] % 8
bs_offs_5 = bs_offs_4 % 4
bs_offs_4 = bs_offs_4 // 4
bs_offs = (
bs_offs_1
+ bs_offs_4 * 2
+ bs_offs_2 * 2 * 2
+ bs_offs_5 * 2 * 2 * 16
+ bs_offs_3 * 2 * 2 * 16 * 4
+ bs_offs_0 * 2 * 16 * scaleN
)
bs_mask1 = (bs_offs_m < M)[:, None] & (bs_offs_n < scaleN)[None, :]
bs_mask2 = (bs_offs_m < scaleM_pad)[:, None] & (bs_offs_n < scaleN_pad)[
None, :
]
bs_e8m0 = tl.where(bs_mask1, bs_e8m0, 127)
tl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask2)
if _USE_INLINE:
try:
from torch.utils.cpp_extension import load_inline
CPP_SRC = r"""
#include <torch/extension.h>
void mxfp4_fused(torch::Tensor A,
torch::Tensor B_shuffle,
torch::Tensor B_scale_sh,
torch::Tensor out);
"""
HIP_SRC = r"""
#include <torch/extension.h>
#include <stdexcept>
void mxfp4_fused(torch::Tensor A,
torch::Tensor B_shuffle,
torch::Tensor B_scale_sh,
torch::Tensor out) {
TORCH_CHECK(false, "mxfp4_fused HIP kernel not implemented yet");
}
"""
_inline_mod = load_inline(
name="mxfp4_inline",
cpp_sources=[CPP_SRC],
cuda_sources=[HIP_SRC],
functions=["mxfp4_fused"],
verbose=_DEBUG,
extra_cuda_cflags=[f"--offload-arch={_INLINE_ARCH}", "-std=c++20"],
)
_INLINE_AVAILABLE = True
except Exception as e:
_INLINE_AVAILABLE = False
if _DEBUG:
print(f"[mxfp4] inline compile failed: {e}")
def _quantize_a(A: torch.Tensor):
if _CACHE_AQ:
key = (A.data_ptr(), getattr(A, "_version", None))
entry = _AQ_REUSE_CACHE.get(key)
if entry is not None:
aref, A_q, A_scale_sh = entry
if aref() is A:
return A_q, A_scale_sh
# stale entry
_AQ_REUSE_CACHE.pop(key, None)
if not _FUSED_SHUFFLE:
A_fp4, A_scale = _quant(A)
A_q = A_fp4.view(_FP4X2)
A_scale_sh = _shuffle(A_scale).view(_FP8E8M0)
if _CACHE_AQ:
key = (A.data_ptr(), getattr(A, "_version", None))
if key in _AQ_REUSE_CACHE:
_AQ_REUSE_ORDER.remove(key)
_AQ_REUSE_CACHE[key] = (weakref.ref(A), A_q, A_scale_sh)
_AQ_REUSE_ORDER.append(key)
if len(_AQ_REUSE_ORDER) > _AQ_REUSE_MAX:
old = _AQ_REUSE_ORDER.pop(0)
_AQ_REUSE_CACHE.pop(old, None)
return A_q, A_scale_sh
M, K = A.shape
scaleN_valid = (K + 31) // 32
scaleN_pad = (scaleN_valid + 7) // 8 * 8
scaleM_pad = (M + 255) // 256 * 256
key = (M, K, A.device)
entry = _AQ_CACHE.get(key)
if entry is None:
A_fp4_buf = torch.empty((M, K // 2), dtype=torch.uint8, device=A.device)
A_scale_sh_buf = torch.empty(
(scaleM_pad, scaleN_pad), dtype=torch.uint8, device=A.device
)
entry = (A_fp4_buf, A_scale_sh_buf, scaleN_valid, scaleM_pad, scaleN_pad)
_AQ_CACHE[key] = entry
else:
A_fp4_buf, A_scale_sh_buf, _, _, _ = entry
# Match aiter.ops.triton.quant.dynamic_mxfp4_quant heuristics
if M <= 32:
NUM_ITER = 1
BLOCK_SIZE_M = triton.next_power_of_2(M)
BLOCK_SIZE_N = 32
NUM_WARPS = 1
NUM_STAGES = 1
else:
NUM_ITER = 4
BLOCK_SIZE_M = 64
BLOCK_SIZE_N = 64
NUM_WARPS = 4
NUM_STAGES = 2
if K <= 16384:
BLOCK_SIZE_M = 32
BLOCK_SIZE_N = 128
if K <= 1024:
NUM_ITER = 1
NUM_STAGES = 1
NUM_WARPS = 4
BLOCK_SIZE_N = min(256, triton.next_power_of_2(K))
BLOCK_SIZE_N = max(32, BLOCK_SIZE_N)
BLOCK_SIZE_M = min(8, triton.next_power_of_2(M))
grid = (
triton.cdiv(M, BLOCK_SIZE_M),
triton.cdiv(K, BLOCK_SIZE_N * NUM_ITER),
)
_dynamic_mxfp4_quant_kernel_shuffled[grid](
A,
A_fp4_buf,
A_scale_sh_buf,
*A.stride(),
*A_fp4_buf.stride(),
*A_scale_sh_buf.stride(),
M=M,
N=K,
scaleN=scaleN_valid,
scaleM_pad=scaleM_pad,
scaleN_pad=scaleN_pad,
BLOCK_SIZE_M=BLOCK_SIZE_M,
BLOCK_SIZE_N=BLOCK_SIZE_N,
NUM_ITER=NUM_ITER,
NUM_STAGES=NUM_STAGES,
MXFP4_QUANT_BLOCK_SIZE=32,
num_warps=NUM_WARPS,
waves_per_eu=0,
num_stages=NUM_STAGES,
)
A_q = A_fp4_buf.view(_FP4X2)
A_scale_sh = A_scale_sh_buf.view(_FP8E8M0)
if _CACHE_AQ:
key = (A.data_ptr(), getattr(A, "_version", None))
if key in _AQ_REUSE_CACHE:
_AQ_REUSE_ORDER.remove(key)
_AQ_REUSE_CACHE[key] = (weakref.ref(A), A_q, A_scale_sh)
_AQ_REUSE_ORDER.append(key)
if len(_AQ_REUSE_ORDER) > _AQ_REUSE_MAX:
old = _AQ_REUSE_ORDER.pop(0)
_AQ_REUSE_CACHE.pop(old, None)
return A_q, A_scale_sh
def _resolve_kernel_candidates():
if _KERNEL_CANDS_ENV.strip().lower() == "auto":
ks = _load_f4gemm_kernels()
if ks:
return ks
alias = {
"32x128": _KERNEL_32,
"192x128": _KERNEL_192,
_KERNEL_32: _KERNEL_32,
_KERNEL_192: _KERNEL_192,
}
out = []
for item in _KERNEL_CANDS_ENV.split(","):
key = item.strip()
if not key:
continue
out.append(alias.get(key, key))
# de-dup while preserving order
seen = set()
uniq = []
for k in out:
if k not in seen:
uniq.append(k)
seen.add(k)
return uniq or [_KERNEL_32, _KERNEL_192]
def _parse_splitk_candidates():
out = []
for item in _SPLITK_CANDS_ENV.split(","):
item = item.strip()
if not item:
continue
try:
out.append(int(item))
except ValueError:
pass
return out or [0]
def _time_gemm_asm(A_q, B_shuffle, A_scale_sh, B_scale_sh, out_buf, kernel, splitK):
# Warmup
_gemm_asm(
A_q, B_shuffle, A_scale_sh, B_scale_sh,
out_buf, kernel,
None, 1.0, 0.0, True,
log2_k_split=splitK,
)
torch.cuda.synchronize()
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
for _ in range(_AUTOTUNE_REPS):
_gemm_asm(
A_q, B_shuffle, A_scale_sh, B_scale_sh,
out_buf, kernel,
None, 1.0, 0.0, True,
log2_k_split=splitK,
)
end.record()
torch.cuda.synchronize()
return start.elapsed_time(end) / _AUTOTUNE_REPS
def _autotune_kernel_splitk(A_q, B_shuffle, A_scale_sh, B_scale_sh, out_buf, kernelName, splitK):
if not (_AUTOTUNE_KERNEL or _AUTOTUNE_SPLITK):
return kernelName, splitK
kernels = [kernelName]
if _AUTOTUNE_KERNEL:
kernels = _resolve_kernel_candidates()
splitks = [splitK]
if _AUTOTUNE_SPLITK:
splitks = _parse_splitk_candidates()
best = (kernelName, splitK, float("inf"))
for kname in kernels:
if "_ZN" not in kname:
continue
for sk in splitks:
try:
t_ms = _time_gemm_asm(
A_q, B_shuffle, A_scale_sh, B_scale_sh,
out_buf, kname, sk,
)
except Exception:
continue
if t_ms < best[2]:
best = (kname, sk, t_ms)
if best[2] < float("inf"):
if _DEBUG:
print(f"[mxfp4] autotune best kernel={best[0]} splitK={best[1]} {best[2]:.4f} ms")
return best[0], best[1]
return kernelName, splitK
def _get_or_create_bufs(M, N, K, device):
key = (M, N, K)
if key in _cache:
return _cache[key]
padded_M = (M + 31) // 32 * 32
out_buf = torch.empty((padded_M, N), dtype=_BF16, device=device)
tuned = _TUNED_MAP.get((M, N, K))
if tuned is not None:
kernelName, splitK = tuned
use_asm = "_ZN" in kernelName
entry = (padded_M, out_buf, use_asm, kernelName, splitK)
_cache[key] = entry
return entry
ck_config = _get_cfg(M, N, K)
splitK = 0
kernelName = ""
use_asm = True # default when no config
if ck_config is not None:
splitK = ck_config.get("splitK", 0) or 0
kernelName = ck_config["kernelName"]
use_asm = "_ZN" in kernelName
if ck_config is None and _AUTO_KERNEL and not _FORCE_KERNEL:
if K >= 2048:
kernelName = _KERNEL_192
else:
kernelName = _KERNEL_32
use_asm = True
if _DEBUG:
print(f"[mxfp4] auto kernel={kernelName} (K={K})")
if _FORCE_KERNEL:
kernelName = _FORCE_KERNEL
use_asm = "_ZN" in kernelName
if _DEBUG:
print(f"[mxfp4] force kernel={kernelName}")
if _FORCE_SPLITK:
try:
splitK = int(_FORCE_SPLITK)
if _DEBUG:
print(f"[mxfp4] force splitK={splitK}")
except ValueError:
if _DEBUG:
print(f"[mxfp4] invalid splitK='{_FORCE_SPLITK}', using {splitK}")
entry = (padded_M, out_buf, use_asm, kernelName, splitK)
_cache[key] = entry
return entry
def _try_get_graph(
M, N, K,
A, B_shuffle, B_scale_sh,
out_buf, use_asm, kernelName, splitK,
):
global _GRAPH_HITS, _GRAPH_MISSES, _GRAPH_FAILS
if not _GRAPH_ENABLED:
return None
key = (M, N, K)
entry = _GRAPH_CACHE.get(key)
if entry is not None:
_GRAPH_HITS += 1
return entry
if key in _GRAPH_BLACKLIST:
return None
try:
A_static = torch.empty(A.shape, dtype=A.dtype, device=A.device)
B_shuffle_static = torch.empty(
B_shuffle.shape, dtype=B_shuffle.dtype, device=B_shuffle.device
)
B_scale_static = torch.empty(
B_scale_sh.shape, dtype=B_scale_sh.dtype, device=B_scale_sh.device
)
A_static.copy_(A)
B_shuffle_static.copy_(B_shuffle)
B_scale_static.copy_(B_scale_sh)
g = torch.cuda.CUDAGraph()
torch.cuda.synchronize()
with torch.cuda.graph(g):
A_q, A_scale_sh = _quantize_a(A_static)
if use_asm:
_gemm_asm(
A_q, B_shuffle_static, A_scale_sh, B_scale_static,
out_buf, kernelName,
None, 1.0, 0.0, True,
log2_k_split=splitK,
)
else:
_gemm_blk(
A_q, B_shuffle_static, A_scale_sh, B_scale_static,
out_buf, splitK,
)
entry = {
"graph": g,
"A_static": A_static,
"B_shuffle_static": B_shuffle_static,
"B_scale_static": B_scale_static,
"out_buf": out_buf,
"b_ptr": B_shuffle.data_ptr(),
"bs_ptr": B_scale_sh.data_ptr(),
# keep intermediates alive for stable graph addresses
"A_q": A_q,
"A_scale_sh": A_scale_sh,
}
_GRAPH_CACHE[key] = entry
_GRAPH_MISSES += 1
return entry
except Exception:
_GRAPH_FAILS += 1
_GRAPH_BLACKLIST.add(key)
return None
@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
global _warmed, _INLINE_AVAILABLE, _INLINE_WARNED
A, B, B_q, B_shuffle, B_scale_sh = data
if not A.is_contiguous():
A = A.contiguous()
if not B_shuffle.is_contiguous():
B_shuffle = B_shuffle.contiguous()
if not B_scale_sh.is_contiguous():
B_scale_sh = B_scale_sh.contiguous()
M, K = A.shape
N = B.shape[0]
# Warm up JIT on first call (builds .so files)
if not _warmed:
A_q, A_scale_sh = _quantize_a(A)
_warmed = True
return aiter.gemm_a4w4(
A_q, B_shuffle, A_scale_sh, B_scale_sh,
dtype=_BF16, bpreshuffle=True,
)
padded_M, out_buf, use_asm, kernelName, splitK = \
_get_or_create_bufs(M, N, K, A.device)
should_autotune = (_AUTOTUNE_KERNEL or _AUTOTUNE_SPLITK) and (
_AUTOTUNE_ALWAYS or (M, N, K) not in _TUNED_MAP
)
if should_autotune:
A_q, A_scale_sh = _quantize_a(A)
kernelName, splitK = _autotune_kernel_splitk(
A_q, B_shuffle, A_scale_sh, B_scale_sh,
out_buf, kernelName, splitK,
)
use_asm = "_ZN" in kernelName
_TUNED_MAP[(M, N, K)] = (kernelName, splitK)
_cache[(M, N, K)] = (padded_M, out_buf, use_asm, kernelName, splitK)
if _PRINT_TUNED_MAP and (M, N, K) not in _TUNED_MAP_PRINTED:
_TUNED_MAP_PRINTED.add((M, N, K))
print(f"[mxfp4] tuned {M},{N},{K} -> kernel={kernelName} splitK={splitK}")
if _INLINE_AVAILABLE and _INLINE_EXEC:
try:
_inline_mod.mxfp4_fused(A, B_shuffle, B_scale_sh, out_buf)
return out_buf[:M]
except Exception as e:
_INLINE_AVAILABLE = False
if _DEBUG:
print(f"[mxfp4] inline exec failed, fallback: {e}")
elif _INLINE_AVAILABLE and _DEBUG and not _INLINE_WARNED:
_INLINE_WARNED = True
print("[mxfp4] inline kernel compiled but not executed (set MXFP4_INLINE_EXEC=1)")
graph_entry = _try_get_graph(
M, N, K,
A, B_shuffle, B_scale_sh,
out_buf, use_asm, kernelName, splitK,
)
if graph_entry is not None:
graph_entry["A_static"].copy_(A)
if graph_entry["b_ptr"] != B_shuffle.data_ptr():
graph_entry["B_shuffle_static"].copy_(B_shuffle)
graph_entry["b_ptr"] = B_shuffle.data_ptr()
if graph_entry["bs_ptr"] != B_scale_sh.data_ptr():
graph_entry["B_scale_static"].copy_(B_scale_sh)
graph_entry["bs_ptr"] = B_scale_sh.data_ptr()
graph_entry["graph"].replay()
return out_buf[:M]
A_q, A_scale_sh = _quantize_a(A)
# Pass unpadded A_q and A_scale_sh -- the ASM kernel handles M-padding internally
if use_asm:
_gemm_asm(
A_q, B_shuffle, A_scale_sh, B_scale_sh,
out_buf, kernelName,
None, 1.0, 0.0, True,
log2_k_split=splitK,
)
else:
_gemm_blk(
A_q, B_shuffle, A_scale_sh, B_scale_sh,
out_buf, splitK,
)
return out_buf[:M]
scrolls · 726 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON