submission 711568
vuxml · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 326 lines, June 9 Researcher Reciprocity License v1.0.
submission_v10d_hybrid.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-711568?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:59c6f64defae9d216e84a9e6be65a0cc63cc0017673c7e2535c27aa140b1d2c1
license declaredunknown
license concludedunknown
authorsvuxml
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
v10d: hybrid fused vs {quant-kernel + fp4-GEMM}. Autotune (L2-cold)num-warps = 4
BM=QBM, BK=QBK, num_warps=4)split-k
SPLIT_K: tl.constexpr, EVEN_N: tl.constexpr, PREQUANT: tl.constexpr,tile-k = 16
QBM, QBK = 16, min(256, k)Kernel source
submission_v10d_hybrid.py326 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
v10d: hybrid fused vs {quant-kernel + fp4-GEMM}. Autotune (L2-cold)
picks per shape. All launches via HIPLauncher (skip JITFunction.run).
Fused is optimal for small m·n (quant cost amortized). Prequant is
optimal for large m·n (A quantized once, not per N-tile).
"""
import os, sys
os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
import warnings; warnings.filterwarnings("ignore")
import torch
import triton
import triton.language as tl
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.triton._triton_kernels.quant.quant import _mxfp4_quant_op
_L = lambda *a: print(*a, file=sys.stderr, flush=True)
_DRV = triton.runtime.driver.active
_GCS = getattr(_DRV, "get_current_" + chr(115) + "tream")
# e8m0_shuffle forward flat idx: view(sm//32,2,16, sn//8,2,4).permute(0,3,5,2,4,1)
@triton.jit
def _sh_row(r, sn8):
return (r // 32) * (sn8 * 256) + (r % 16) * 4 + (r // 16) % 2
@triton.jit
def _sh_col(c):
return (c // 8) * 256 + (c % 4) * 64 + (c // 4) % 2 * 2
@triton.jit
def _gemm_k(
A, Asc, Bq, Bsc, C,
M, N, K, sA_m, sAsc_m, sBq_n, sC_k, sC_m, sn8,
BM: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr,
SPLIT_K: tl.constexpr, EVEN_N: tl.constexpr, PREQUANT: tl.constexpr,
):
pid = tl.program_id(0)
num_n = tl.cdiv(N, BN)
num_mn = tl.cdiv(M, BM) * num_n
pid_k = pid // num_mn
pid_mn = pid % num_mn
pid_m = pid_mn // num_n
pid_n = pid_mn % num_n
offs_m = pid_m * BM + tl.arange(0, BM)
offs_n = pid_n * BN + tl.arange(0, BN)
offs_n64 = offs_n.to(tl.int64)
mask_m = offs_m < M
mask_n = offs_n < N
rk = tl.arange(0, BK)
rk2 = tl.arange(0, BK // 2)
rk32 = tl.arange(0, BK // 32)
k_per = tl.cdiv(tl.cdiv(K, BK), SPLIT_K) * BK
k_lo = pid_k * k_per
k_hi = min(k_lo + k_per, K)
bq_ptrs = Bq + offs_n64[:, None] * sBq_n + (k_lo // 2 + rk2)[None, :]
bsc_row = _sh_row(offs_n64, sn8)
acc = tl.zeros((BM, BN), dtype=tl.float32)
if PREQUANT:
a_ptrs = A + offs_m[:, None].to(tl.int64) * sA_m + (k_lo // 2 + rk2)[None, :]
asc_ptrs = Asc + offs_m[:, None].to(tl.int64) * sAsc_m + (k_lo // 32 + rk32)[None, :]
else:
a_ptrs = A + offs_m[:, None].to(tl.int64) * sA_m + (k_lo + rk)[None, :]
for k in tl.range(k_lo, k_hi, BK):
if PREQUANT:
a_fp4 = tl.load(a_ptrs, mask=mask_m[:, None], other=0)
a_sc = tl.load(asc_ptrs, mask=mask_m[:, None], other=0)
else:
a_bf = tl.load(a_ptrs, mask=mask_m[:, None], other=0.0)
a_fp4, a_sc = _mxfp4_quant_op(a_bf.to(tl.float32), BK, BM, 32)
if EVEN_N:
b_fp4_t = tl.load(bq_ptrs)
else:
b_fp4_t = tl.load(bq_ptrs, mask=mask_n[:, None], other=0)
bsc_col = _sh_col(k // 32 + rk32)
if EVEN_N:
b_sc = tl.load(Bsc + bsc_row[:, None] + bsc_col[None, :])
else:
b_sc = tl.load(Bsc + bsc_row[:, None] + bsc_col[None, :],
mask=mask_n[:, None], other=0)
acc = tl.dot_scaled(a_fp4, a_sc, "e2m1",
tl.trans(b_fp4_t), b_sc, "e2m1", acc)
if PREQUANT:
a_ptrs += BK // 2
asc_ptrs += BK // 32
else:
a_ptrs += BK
bq_ptrs += BK // 2
c_off = (pid_k * sC_k + offs_m[:, None].to(tl.int64) * sC_m
+ offs_n[None, :])
cmask = mask_m[:, None] & mask_n[None, :]
if SPLIT_K == 1:
tl.store(C + c_off, acc.to(tl.bfloat16), mask=cmask)
else:
tl.store(C + c_off, acc, mask=cmask)
@triton.jit
def _reduce_k(W, C, SK, M, N, sW_k, sW_m, sC_m,
BLK: tl.constexpr, SKC: tl.constexpr):
pid = tl.program_id(0)
off = pid * BLK + tl.arange(0, BLK)
om = off // N; on = off % N; mask = om < M
base = om.to(tl.int64) * sW_m + on
s = tl.zeros((BLK,), dtype=tl.float32)
for i in tl.static_range(SKC):
s += tl.load(W + i * sW_k + base, mask=mask & (i < SK), other=0.0)
tl.store(C + om.to(tl.int64) * sC_m + on, s.to(tl.bfloat16), mask=mask)
@triton.jit
def _quant_a_k(A, Afp4, Asc, M, K, sA_m, sAf_m, sAs_m,
BM: tl.constexpr, BK: tl.constexpr):
pid = tl.program_id(0)
nk = tl.cdiv(K, BK)
pm = pid // nk; pk = pid % nk
offs_m = pm * BM + tl.arange(0, BM)
offs_k = pk * BK + tl.arange(0, BK)
mask_m = offs_m < M
a = tl.load(A + offs_m[:, None].to(tl.int64) * sA_m + offs_k[None, :],
mask=mask_m[:, None], other=0.0).to(tl.float32)
af, asc = _mxfp4_quant_op(a, BK, BM, 32)
tl.store(Afp4 + offs_m[:, None].to(tl.int64) * sAf_m
+ (pk * (BK // 2) + tl.arange(0, BK // 2))[None, :],
af, mask=mask_m[:, None])
tl.store(Asc + offs_m[:, None].to(tl.int64) * sAs_m
+ (pk * (BK // 32) + tl.arange(0, BK // 32))[None, :],
asc, mask=mask_m[:, None])
def _cfgs(m, n, k):
"""(BM, BN, BK, SPLIT_K, nw, nK, ns, PREQUANT).
m≤32: fused only. m≥64: prequant only, BM∈{32,64}."""
out = []
PQs = (False,) if m <= 32 else (True,) if m >= 128 else (False, True)
for PQ in PQs:
BMs = (16,) if not PQ else tuple(b for b in (32, 64, 128) if b <= m)
for BM in BMs:
for BN in (32, 64, 128, 256):
if BN > n: continue
bt = -(-m // BM) * -(-n // BN)
for BK in (256, 512):
if BK > k: continue
nkit = k // BK
SKs = [1]
if bt < 200 and nkit >= 2:
tgt = max(1, 256 // bt)
for s in (2, 4, 8):
if s <= nkit and s <= tgt * 2: SKs.append(s)
for SK in SKs:
for nw in (4, 8):
for nK in ((16,) if BM == 16 else (16, 32)):
out.append((BM, BN, BK, SK, nw, nK, 2, PQ))
seen, r = set(), []
for c in out:
if c not in seen: seen.add(c); r.append(c)
return r
_L2FLUSH = torch.empty(512 * 1024 * 1024, dtype=torch.int8, device="cuda")
def _gpu_time_cold(fn, n_iter=6):
for _ in range(2): fn()
torch.cuda.synchronize()
evs = [(torch.cuda.Event(True), torch.cuda.Event(True)) for _ in range(n_iter)]
for e0, e1 in evs:
_L2FLUSH.zero_()
e0.record(); fn(); e1.record()
torch.cuda.synchronize()
return sum(e0.elapsed_time(e1) for e0, e1 in evs) * 1000.0 / n_iter
def _ref(A, B_shuffle, B_scale_sh):
Aq, As = dynamic_mxfp4_quant(A)
As = e8m0_shuffle(As)
return aiter.gemm_a4w4(Aq.view(dtypes.fp4x2), B_shuffle,
As.view(dtypes.fp8_e8m0), B_scale_sh,
dtype=dtypes.bf16, bpreshuffle=True)
_STATE: dict = {}
def _mk_launcher(kernel, grid_x, const_args, **kw):
ck = kernel.warmup(*const_args, grid=(grid_x,), **kw)
ck._init_handles()
return ck.run, ck.function, ck.packed_metadata
def _build(data):
A, B, B_q, B_shuffle, B_scale_sh = data
m, k = A.shape; n = B.shape[0]
sn = B_scale_sh.shape[1]; sn8 = sn // 8
dev = A.device
Bq = B_q.contiguous().view(torch.uint8)
Bsc = B_scale_sh.contiguous().view(torch.uint8).reshape(-1)
C_bf = torch.empty((m, n), dtype=torch.bfloat16, device=dev)
W = torch.empty((8, m, n), dtype=torch.float32, device=dev)
Afp4 = torch.empty((m, k // 2), dtype=torch.uint8, device=dev)
Asc = torch.empty((m, k // 32), dtype=torch.uint8, device=dev)
ref_f = _ref(A, B_shuffle, B_scale_sh).float()
mag = ref_f.abs().mean().item() + 1e-9
try:
lq = _GCS(None)
except Exception:
lq = _GCS(torch.cuda.current_device())
QBM, QBK = 16, min(256, k)
q_gx = triton.cdiv(m, QBM) * (k // QBK)
q_run, q_fun, q_pmeta = _mk_launcher(
_quant_a_k, q_gx,
(A, Afp4, Asc, m, k, k, k // 2, k // 32),
BM=QBM, BK=QBK, num_warps=4)
r_gx = triton.cdiv(m * n, 256)
r_run, r_fun, r_pmeta = _mk_launcher(
_reduce_k, r_gx,
(W, C_bf, 8, m, n, m * n, n, n),
BLK=256, SKC=8, num_warps=4)
def _do_quant(A_in):
q_run(q_gx, 1, 1, lq, q_fun, q_pmeta, None, None, None,
A_in, Afp4, Asc, m, k, k, k // 2, k // 32, QBM, QBK)
def _do_reduce(SK):
r_run(r_gx, 1, 1, lq, r_fun, r_pmeta, None, None, None,
W, C_bf, SK, m, n, m * n, n, n, 256, 8)
cfgs = _cfgs(m, n, k)
_L(f"\n[v10d m={m} n={n} k={k}] {len(cfgs)} cfgs")
best, best_t, best_go = None, float("inf"), None
for cfg in cfgs:
BM, BN, BK, SK, nw, nK, ns, PQ = cfg
C_out = W if SK > 1 else C_bf
sC_k = W.stride(0) if SK > 1 else 0
sC_m = W.stride(1) if SK > 1 else n
sAm = (k // 2) if PQ else k
even_n = (n % BN == 0)
gx = triton.cdiv(m, BM) * triton.cdiv(n, BN) * SK
A_in = Afp4 if PQ else A
try:
run, fun, pmeta = _mk_launcher(
_gemm_k, gx,
(A_in, Asc, Bq, Bsc, C_out, m, n, k,
sAm, k // 32, k // 2, sC_k, sC_m, sn8),
BM=BM, BN=BN, BK=BK, SPLIT_K=SK, EVEN_N=even_n,
PREQUANT=PQ, num_warps=nw, num_stages=ns,
matrix_instr_nonkdim=nK, waves_per_eu=0)
except Exception as e:
if best is None:
_L(f" {cfg}: COMPILE {type(e).__name__}: {str(e)[:100]}")
continue
cargs = (BM, BN, BK, SK, even_n, PQ)
def _go(_A=A, _Bq=Bq, _Bsc=Bsc, _run=run, _fun=fun,
_pmeta=pmeta, _cargs=cargs, _gx=gx, _PQ=PQ, _SK=SK,
_C=C_out, _sAm=sAm, _sCk=sC_k, _sCm=sC_m):
if _PQ:
_do_quant(_A)
_run(_gx, 1, 1, lq, _fun, _pmeta, None, None, None,
Afp4, Asc, _Bq, _Bsc, _C, m, n, k,
_sAm, k // 32, k // 2, _sCk, _sCm, sn8, *_cargs)
else:
_run(_gx, 1, 1, lq, _fun, _pmeta, None, None, None,
_A, Asc, _Bq, _Bsc, _C, m, n, k,
_sAm, k // 32, k // 2, _sCk, _sCm, sn8, *_cargs)
if _SK > 1:
_do_reduce(_SK)
return C_bf
try:
out = _go()
torch.cuda.synchronize()
err = ((out.float() - ref_f).abs().mean() / mag).item()
if err > 5e-3:
if best is None: _L(f" {cfg}: ERR {err:.2%}")
continue
t = _gpu_time_cold(_go)
if t < best_t:
best_t, best, best_go = t, cfg, _go
_L(f" {cfg}: {t:.2f}us grid={gx} *")
except Exception as e:
if best is None:
_L(f" {cfg}: RUN {type(e).__name__}: {str(e)[:100]}")
if best is None:
_L(f" → fallback")
return {"hot": None}
_L(f" → best={best} @ {best_t:.2f}us")
return {"hot": best_go}
def custom_kernel(data):
A = data[0]; Bq = data[2]
key = (A.shape[0], Bq.shape[0], A.shape[1])
S = _STATE.get(key)
if S is None:
S = _build(data); _STATE[key] = S
hot = S["hot"]
if hot is None:
return _ref(A, data[3], data[4])
return hot(A, Bq, data[4])
scrolls · 326 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 705642.
⋯ 151 unchanged linesfor PQ in PQs:BMs = (16,) if not PQ else tuple(b for b in (32, 64, 128) if b <= m)for BM in BMs:- for BN in (64, 128, 256):+ for BN in (32, 64, 128, 256):if BN > n: continuebt = -(-m // BM) * -(-n // BN)for BK in (256, 512):⋯ 148 unchanged linesif best is None:_L(f" → fallback")- return {"fallback": True}+ return {"hot": None}_L(f" → best={best} @ {best_t:.2f}us")- return {"fallback": False, "hot": best_go}+ return {"hot": best_go}def custom_kernel(data):- A, B, B_q, B_shuffle, B_scale_sh = data- m, k = A.shape; n = B.shape[0]- key = (m, n, k)+ A = data[0]; Bq = data[2]+ key = (A.shape[0], Bq.shape[0], A.shape[1])S = _STATE.get(key)if S is None:S = _build(data); _STATE[key] = S- if S["fallback"]:- return _ref(A, B_shuffle, B_scale_sh)- return S["hot"](A, B_q.view(torch.uint8),- B_scale_sh.view(torch.uint8).reshape(-1))+ hot = S["hot"]+ if hot is None:+ return _ref(A, data[3], data[4])+ return hot(A, Bq, data[4])
scrolls · 38 diff lines total
Best evidence level for this revision: reported
JSON