submission 748551
Chivier · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 356 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-748551?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:99e23ceff36841297414fc91d577e996422ee8bcad7539a4a2c01d23f1dfe18c
license declaredunknown
license concludedunknown
authorsChivier
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MXFP4 GEMM V276 — F32 CVT variant + int-ptr replay + minimal hot path.tile-m = 16
C.stride(0), C.stride(1), KS=ks, BM=16, BN=128,tile-n = 64
(64, 7168, 2048): (16, 128, 512, 4, 2, 8, 1), # V234: 12.0µs (BN=64 regressed to 13.5)Kernel source
submission.py356 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
MXFP4 GEMM V276 — F32 CVT variant + int-ptr replay + minimal hot path.
CVT change: V_CVT_SCALEF32_PK_FP4_F32 (opcode 573) takes two f32 inputs directly.
Eliminates: bf16 conversion + uint16 bitcast + split + pack = ~6 ops per pair.
F32 variant: f32_to_fp4_scale(S0.f32, scale) — full f32 precision, no bf16 intermediate.
Also: pre-cache bs_sh view, minimal Python hot path.
"""
import os
os.environ['HIP_FORCE_DEV_KERNARG'] = '1'
import torch
import triton
import triton.language as tl
from task import input_t, output_t
# ─── Inlined _mxfp4_quant_op with bit-ops replacing log2/exp2 ───
@triton.jit
def _mxfp4_quant_op(
x,
BLOCK_SIZE_N,
BLOCK_SIZE_M,
MXFP4_QUANT_BLOCK_SIZE,
):
"""F32 CVT quant: software f32 scale + V_CVT_SCALEF32_PK_FP4_F32.
F32 variant: takes two separate f32 values + f32 scale → packed FP4.
No bf16 conversion needed — full f32 precision throughout.
ISA opcode 573: f32_to_fp4_scale(S0.f32, scale) + f32_to_fp4_scale(S1.f32, scale)
"""
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
HQB: tl.constexpr = MXFP4_QUANT_BLOCK_SIZE // 2
x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
# ── Step 1: amax + E8M0 scale ──
amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
amax = amax.to(tl.int32, bitcast=True)
amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax_exp = (amax >> 23) & 0xFF
scale_e8m0_unbiased = amax_exp.to(tl.int32) - 129
scale_e8m0_unbiased = tl.maximum(scale_e8m0_unbiased, -127)
scale_e8m0_unbiased = tl.minimum(scale_e8m0_unbiased, 127)
bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
# ── Step 2: Inverse scale ──
inv_exp = (256 - amax_exp).to(tl.int32)
inv_exp = tl.maximum(inv_exp, 0)
inv_exp = tl.minimum(inv_exp, 254)
quant_scale = (inv_exp.to(tl.uint32) << 23).to(tl.float32, bitcast=True)
amax_f = amax.to(tl.float32, bitcast=True)
quant_scale = tl.where(amax_f > 0, quant_scale, tl.zeros_like(quant_scale))
# ── Step 3: Scale in f32 ──
qx = x * quant_scale # [BM, NQB, QBS] f32
# ── Step 4: F32 CVT — skip bf16 conversion entirely ──
# Split into even/odd f32 values for CVT_PK_FP4_F32(S0, S1, scale)
qx_pairs = qx.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, HQB, 2)
qx_even, qx_odd = tl.split(qx_pairs) # each [BM, NQB, HQB, 1]
qx_even = qx_even.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, HQB)
qx_odd = qx_odd.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, HQB)
one_f32 = tl.full([BLOCK_SIZE_M, NUM_QUANT_BLOCKS, HQB], 1.0, dtype=tl.float32)
# V_CVT_SCALEF32_PK_FP4_F32: vdst, S0(f32), S1(f32), S2(scale_f32)
# Output: byte at OPSEL position = {fp4(S1), fp4(S0)}
fp4_u32 = tl.inline_asm_elementwise(
"v_cvt_scalef32_pk_fp4_f32 $0, $1, $2, $3",
"=v,v,v,v",
args=[qx_even, qx_odd, one_f32],
dtype=tl.uint32,
is_pure=True,
pack=1,
)
x_fp4 = (fp4_u32 & 0xFF).to(tl.uint8)
x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
@triton.jit
def _remap_xcd(pid, GRID, NX: tl.constexpr):
ppx = (GRID + NX - 1) // NX
tx = GRID % NX
if tx == 0: tx = NX
x = pid % NX
lp = pid // NX
if x < tx: return x * ppx + lp
else: return tx * ppx + (x - tx) * (ppx - 1) + lp
@triton.jit
def _fused(
A, Bq, Bs_sh, C,
M, N, K,
sa0, sa1, sbq0, sbq1, sc0, sc1,
BM: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr,
GM: tl.constexpr, NX: tl.constexpr,
EVEN_K: tl.constexpr,
):
pid = tl.program_id(0)
nm = tl.cdiv(M, BM); nn = tl.cdiv(N, BN)
pid = _remap_xcd(pid, nm * nn, NX)
gn = GM * nn; gid = pid // gn; fm = gid * GM
gsm = min(nm - fm, GM); pm = fm + (pid % gsm)
pn = (pid % gn) // gsm
if pm >= nm or pn >= nn: return
om = pm * BM + tl.arange(0, BM)
on = pn * BN + tl.arange(0, BN)
acc = tl.zeros((BM, BN), dtype=tl.float32)
HK: tl.constexpr = BK // 2
SK: tl.constexpr = BK // 32
sh_n_base = (on // 32) * K + (on % 16) * 4 + (on % 32) // 16
for ks in range(0, K, BK):
ok = ks + tl.arange(0, BK)
if EVEN_K:
a = tl.load(A + om[:, None] * sa0 + ok[None, :] * sa1)
else:
a = tl.load(A + om[:, None] * sa0 + ok[None, :] * sa1,
mask=(om[:, None] < M) & (ok[None, :] < K), other=0.0)
aq, asc = _mxfp4_quant_op(a.to(tl.float32), BK, BM, 32)
kp = ks // 2 + tl.arange(0, HK)
if EVEN_K:
b = tl.load(Bq + kp[:, None] * sbq1 + on[None, :] * sbq0,
cache_modifier=".cg")
else:
b = tl.load(Bq + kp[:, None] * sbq1 + on[None, :] * sbq0,
mask=on[None, :] < N, other=0,
cache_modifier=".cg")
ksc = ks // 32 + tl.arange(0, SK)
sh_k_part = (ksc // 8) * 256 + (ksc % 4) * 64 + ((ksc % 8) // 4) * 2
if EVEN_K:
bs = tl.load(Bs_sh + sh_n_base[:, None] + sh_k_part[None, :],
cache_modifier=".cg")
else:
bs = tl.load(Bs_sh + sh_n_base[:, None] + sh_k_part[None, :],
mask=on[:, None] < N, other=127,
cache_modifier=".cg")
acc = tl.dot_scaled(aq, asc, "e2m1", b, bs, "e2m1", acc=acc, out_dtype=tl.float32)
cp = C + om[:, None] * sc0 + on[None, :] * sc1
if EVEN_K:
tl.store(cp, acc.to(tl.bfloat16))
else:
tl.store(cp, acc.to(tl.bfloat16), mask=(om[:, None] < M) & (on[None, :] < N))
# ─── Split-K kernel for Case 2 (inherited from V39) ───
@triton.jit
def _fused_sk(
A, Bq, Bs_sh, P,
M, N, K,
sa0, sa1, sbq0, sbq1, sp0, sp1, sp2,
BM: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr,
GM: tl.constexpr, NX: tl.constexpr, KS: tl.constexpr,
EVEN_K: tl.constexpr,
):
pid = tl.program_id(0)
nm = tl.cdiv(M, BM); nn = tl.cdiv(N, BN)
GRID = nm * nn * KS
pid = _remap_xcd(pid, GRID, NX)
pk = pid % KS; pmn = pid // KS
pm = pmn // nn; pn = pmn % nn
if pm >= nm or pn >= nn: return
om = pm * BM + tl.arange(0, BM)
on = pn * BN + tl.arange(0, BN)
acc = tl.zeros((BM, BN), dtype=tl.float32)
HK: tl.constexpr = BK // 2
SK: tl.constexpr = BK // 32
sh_n_base = (on // 32) * K + (on % 16) * 4 + (on % 32) // 16
niters = K // (BK * KS)
for ki in range(niters):
ks = (ki * KS + pk) * BK
ok = ks + tl.arange(0, BK)
if EVEN_K:
a = tl.load(A + om[:, None] * sa0 + ok[None, :] * sa1)
else:
a = tl.load(A + om[:, None] * sa0 + ok[None, :] * sa1,
mask=(om[:, None] < M) & (ok[None, :] < K), other=0.0)
aq, asc = _mxfp4_quant_op(a.to(tl.float32), BK, BM, 32)
kp = ks // 2 + tl.arange(0, HK)
if EVEN_K:
b = tl.load(Bq + kp[:, None] * sbq1 + on[None, :] * sbq0,
cache_modifier=".cg")
else:
b = tl.load(Bq + kp[:, None] * sbq1 + on[None, :] * sbq0,
mask=on[None, :] < N, other=0,
cache_modifier=".cg")
ksc = ks // 32 + tl.arange(0, SK)
sh_k_part = (ksc // 8) * 256 + (ksc % 4) * 64 + ((ksc % 8) // 4) * 2
if EVEN_K:
bs = tl.load(Bs_sh + sh_n_base[:, None] + sh_k_part[None, :],
cache_modifier=".cg")
else:
bs = tl.load(Bs_sh + sh_n_base[:, None] + sh_k_part[None, :],
mask=on[:, None] < N, other=127,
cache_modifier=".cg")
acc = tl.dot_scaled(aq, asc, "e2m1", b, bs, "e2m1", acc=acc, out_dtype=tl.float32)
pp = P + pk * sp0 + om[:, None] * sp1 + on[None, :] * sp2
if EVEN_K:
tl.store(pp, acc)
else:
tl.store(pp, acc, mask=(om[:, None] < M) & (on[None, :] < N))
@triton.jit
def _reduce(P, C, M, N, sp0, sp1, sp2, sc0, sc1,
KS: tl.constexpr, BM: tl.constexpr, BN: tl.constexpr):
pm = tl.program_id(0); pn = tl.program_id(1)
om = pm * BM + tl.arange(0, BM)
on = pn * BN + tl.arange(0, BN)
mask = (om[:, None] < M) & (on[None, :] < N)
acc = tl.zeros((BM, BN), dtype=tl.float32)
for s in range(KS):
acc += tl.load(P + s * sp0 + om[:, None] * sp1 + on[None, :] * sp2, mask=mask, other=0.0)
tl.store(C + om[:, None] * sc0 + on[None, :] * sc1, acc.to(tl.bfloat16), mask=mask)
# ─── Per-shape configs: (BM, BN, BK, nw, ns, GM, KS) ───
_CFG = {
# K=512: BK=512 from V200 (single K-iter, proven 6.9-7.1µs)
(4, 2880, 512): (16, 16, 512, 4, 2, 1, 1), # V200: 6.90µs
(32, 4096, 512): (16, 64, 512, 4, 2, 1, 1), # V200: 7.13µs
(32, 2880, 512): (16, 64, 512, 4, 2, 1, 1), # V200: 7.14µs
# Large K: V278 best-of configs
(16, 2112, 7168): (16, 64, 512, 4, 2, 8, 7), # V277: 10.2µs!! (was V39: 14.1µs, -32%)
(64, 7168, 2048): (16, 128, 512, 4, 2, 8, 1), # V234: 12.0µs (BN=64 regressed to 13.5)
(256, 3072, 1536): (16, 128, 256, 4, 3, 8, 1), # V234: 14.0µs (BN=64 regressed to 15.0)
}
_out = {}
_bq = {}
_replay = {} # (m,n,k) → (orig_launch, args_list, ptr_indices)
_ncall = {}
_bs_cache = {}
def _get_bq(bq_raw):
k = id(bq_raw)
c = _bq.get(k)
if c is not None and c[0] is bq_raw: return c[1]
r = bq_raw.view(torch.uint8)
_bq[k] = (bq_raw, r)
return r
def _get_bs(bs_raw):
k = id(bs_raw)
c = _bs_cache.get(k)
if c is not None and c[0] is bs_raw: return c[1]
r = bs_raw.view(torch.uint8)
_bs_cache[k] = (bs_raw, r)
return r
def _get_out(m, n, dev):
k = (m, n)
c = _out.get(k)
if c is not None: return c
c = torch.empty((m, n), dtype=torch.bfloat16, device=dev)
_out[k] = c
return c
def _triton_call(A, bq, bs_sh, C, m, n, k, BM, BN, BK, nw, ns, gm, even_k, grid):
_fused[(grid,)](
A, bq, bs_sh, C, m, n, k,
A.stride(0), A.stride(1), bq.stride(0), bq.stride(1),
C.stride(0), C.stride(1),
BM=BM, BN=BN, BK=BK, GM=gm, NX=8,
EVEN_K=even_k, num_warps=nw, num_stages=ns,
)
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
m, k = A.shape; n = B.shape[0]
bq = _get_bq(B_q)
bs_sh = _get_bs(B_scale_sh)
cfg = _CFG.get((m, n, k))
if cfg is None:
cfg = (16, 128, 256, 4, 3, 8, 1)
BM, BN, BK, nw, ns, gm, ks = cfg
even_k = (k % BK == 0) and (m % BM == 0) and (n % BN == 0)
if ks > 1:
grid_mn = triton.cdiv(m, BM) * triton.cdiv(n, BN)
P = torch.empty((ks, m, n), device=A.device, dtype=torch.float32)
_fused_sk[(grid_mn * ks,)](
A, bq, bs_sh, P, m, n, k,
A.stride(0), A.stride(1), bq.stride(0), bq.stride(1),
P.stride(0), P.stride(1), P.stride(2),
BM=BM, BN=BN, BK=BK, GM=gm, NX=8, KS=ks,
EVEN_K=even_k, num_warps=nw, num_stages=ns,
)
C = _get_out(m, n, A.device)
_reduce[(triton.cdiv(m, 16), triton.cdiv(n, 128))](
P, C, m, n, P.stride(0), P.stride(1), P.stride(2),
C.stride(0), C.stride(1), KS=ks, BM=16, BN=128,
)
return C
C = _get_out(m, n, A.device)
grid = triton.cdiv(m, BM) * triton.cdiv(n, BN)
sk = (m, n, k)
# ─── Fast replay: pass data_ptr() ints instead of Tensors ───
# C launcher: PyLong_Check → PyLong_AsULL (0.01µs)
# vs Tensor: getAttr("data_ptr") → Call → AsULL → hipPointerGetAttribute (0.5µs)
rp = _replay.get(sk)
if rp is not None:
orig_fn, tmpl, pidx = rp
a2 = list(tmpl)
# Pass int ptrs (fast path in C launcher)
a2[pidx[0]] = A.data_ptr()
a2[pidx[1]] = bq.data_ptr()
a2[pidx[2]] = bs_sh.data_ptr()
a2[pidx[3]] = C.data_ptr()
a2[1] = grid # update gridX
try:
orig_fn(*a2)
return C
except Exception as e:
import sys
print(f"[V274] replay FAIL: {e}", file=sys.stderr)
del _replay[sk]
cnt = _ncall.get(sk, 0)
_ncall[sk] = cnt + 1
if cnt == 0:
# 1st call: compile + install capture
import sys
a_ptr = A.data_ptr(); bq_ptr = bq.data_ptr()
bs_ptr = bs_sh.data_ptr(); c_ptr = C.data_ptr()
ret = _fused[(grid,)](
A, bq, bs_sh, C, m, n, k,
A.stride(0), A.stride(1), bq.stride(0), bq.stride(1),
C.stride(0), C.stride(1),
BM=BM, BN=BN, BK=BK, GM=gm, NX=8,
EVEN_K=even_k, num_warps=nw, num_stages=ns,
)
if ret is not None and hasattr(ret, 'run') and hasattr(ret.run, 'launch'):
orig = ret.run.launch
_run_obj = ret.run # hold reference to restore later
def _cap(*args, _sk=sk, _orig=orig, _run=_run_obj):
import sys as _sys, torch as _torch
pidx = []
for i, v in enumerate(args):
if isinstance(v, _torch.Tensor):
pidx.append(i)
if len(pidx) >= 4:
_replay[_sk] = (_orig, list(args), pidx[:4])
# Restore original launch to stop re-capturing
_run.launch = _orig
_sys.stderr.write(f"[V274] CAPTURED {_sk}: pidx={pidx[:4]}\n")
return _orig(*args)
ret.run.launch = _cap
return C
# Normal dispatch
_triton_call(A, bq, bs_sh, C, m, n, k, BM, BN, BK, nw, ns, gm, even_k, grid)
return C
scrolls · 356 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