submission 677752
xg · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 348 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-677752?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:1510080b7dc538b6e961927b26195fe55d70c1f0b88acfcba838679ec96a64d4
license declaredunknown
license concludedunknown
authorsxg
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
Fused MXFP4 quant + GEMM: bf16 A -> inline MXFP4 quant -> tl.dot_scaled GEMM -> bf16 C.tile-n = 16
RBM, RBN = 16, 64Kernel source
submission.py348 lines
"""
Fused MXFP4 quant + GEMM: bf16 A -> inline MXFP4 quant -> tl.dot_scaled GEMM -> bf16 C.
Uses only A, B_q, B_shuffle, B_scale_sh (plus B for API); no input cache, no host-side copies.
B layout matches shuffle_weight(16,16) + e8m0_shuffle (aiter.gemm_a4w4).
"""
from task import input_t, output_t
import torch
import triton
import triton.language as tl
SCALE_GROUP_SIZE = 32
@triton.jit
def _b_shuffle_phys_flat(n_l, k_l, KP):
"""Row-major [N, KP] byte index after shuffle_weight(..., layout=(16,16))."""
k_blk = KP // 32
nb = n_l // 16
ni = n_l % 16
kb = k_l // 32
rem = k_l % 32
sub = rem // 16
ki = rem % 16
return (((nb * k_blk + kb) * 2 + sub) * 16 + ni) * 16 + ki
@triton.jit
def _scale_shuffle_phys_flat(ml, nl, SN):
"""Linear index into row-major padded tensor from e8m0_shuffle (aiter fp4_utils)."""
s1 = SN // 8
a = ml // 32
rem = ml % 32
b = rem // 16
c = rem % 16
d = nl // 8
rem2 = nl % 8
e = rem2 // 4
f = rem2 % 4
return (((a * s1 + d) * 4 + f) * 16 + c) * 4 + e * 2 + b
@triton.jit
def _mxfp4_quant_inline(x, BSK: tl.constexpr, BSM: tl.constexpr, SGS: tl.constexpr):
NQB: tl.constexpr = BSK // SGS
x = x.reshape(BSM, NQB, SGS)
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 = amax.to(tl.float32, bitcast=True)
su = tl.log2(amax).floor() - 2
su = tl.clamp(su, min=-127, max=127)
bs = su.to(tl.uint8) + 127
qx = x * tl.exp2(-su)
qx_u32 = qx.to(tl.uint32, bitcast=True)
s = qx_u32 & 0x80000000
qx_u32 = qx_u32 ^ s
qf = qx_u32.to(tl.float32, bitcast=True)
sat = qf >= 6
den = (not sat) & (qf < 1)
nor = not (sat | den)
de: tl.constexpr = ((127 - 1) + (23 - 1) + 1) << 23
df: tl.constexpr = tl.cast(de, tl.float32, bitcast=True)
dx = qf + df
dx = dx.to(tl.int32, bitcast=True)
dx -= de
dx = dx.to(tl.uint8)
nx = qx_u32.to(tl.int32, bitcast=True)
mo = (nx >> 22) & 1
val_add: tl.constexpr = ((1 - 127) << 23) + (1 << 21) - 1
nx += val_add
nx += mo
nx = nx >> 22
nx = nx.to(tl.uint8)
v = tl.full(qx_u32.type.get_block_shapes(), 0x7, dtype=tl.uint8)
v = tl.where(nor, nx, v)
v = tl.where(den, dx, v)
sl = s >> 28
sl = sl.to(tl.uint8)
v = v | sl
v = tl.reshape(v, [BSM, NQB, SGS // 2, 2])
ev, od = tl.split(v)
fp4 = ev | (od << 4)
fp4 = fp4.reshape(BSM, BSK // 2)
return fp4, bs.reshape(BSM, NQB)
@triton.jit
def _remap_xcd(pid, GM, NX: tl.constexpr = 8):
ppx = (GM + NX - 1) // NX
tx = GM % NX
tx = NX if tx == 0 else tx
xcd = pid % NX
lp = pid // NX
if xcd < tx:
pid = xcd * ppx + lp
else:
pid = tx * ppx + (xcd - tx) * (ppx - 1) + lp
return pid
@triton.jit
def _pgrid(pid, npm, npn, GSM: tl.constexpr = 1):
if GSM == 1:
pm = pid // npn
pn = pid % npn
else:
npig = GSM * npn
gi = pid // npig
fpm = gi * GSM
gsm = min(npm - fpm, GSM)
tl.assume(gsm >= 0)
pm = fpm + (pid % gsm)
pn = (pid % npig) // gsm
return pm, pn
@triton.heuristics({
"EVEN_K": lambda a: (a["KP"] % (a["BSK"] // 2) == 0)
and (a["SPBS"] % a["BSK"] == 0) and (a["KP"] % (a["SPBS"] // 2) == 0),
})
@triton.jit
def _fqg_kernel(
a_ptr, b_ptr, c_ptr, bs_ptr,
M, N, KP, SN, NUM_SG,
sa0, sa1, sb0, sb1, sbs0, sbs1, sck, scm, scn,
BSM: tl.constexpr, BSN: tl.constexpr, BSK: tl.constexpr,
GSM: tl.constexpr, NKS: tl.constexpr, SPBS: tl.constexpr,
EVEN_K: tl.constexpr,
num_warps: tl.constexpr, num_stages: tl.constexpr,
waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
):
tl.assume(sa0 > 0)
tl.assume(sa1 > 0)
tl.assume(sb0 > 0)
tl.assume(sb1 > 0)
tl.assume(sbs0 > 0)
tl.assume(sbs1 > 0)
tl.assume(scm > 0)
tl.assume(scn > 0)
GMN = tl.cdiv(M, BSM) * tl.cdiv(N, BSN)
SGS: tl.constexpr = 32
pu = tl.program_id(0)
pu = _remap_xcd(pu, GMN * NKS, NX=8)
pk = pu % NKS
p = pu // NKS
npm = tl.cdiv(M, BSM)
npn = tl.cdiv(N, BSN)
if NKS == 1:
pm, pn = _pgrid(p, npm, npn, GSM=GSM)
else:
pm = p // npn
pn = p % npn
tl.assume(pm >= 0)
tl.assume(pn >= 0)
tl.assume(pk >= 0)
if (pk * SPBS // 2) < KP:
nki = tl.cdiv(SPBS // 2, BSK // 2)
om = (pm * BSM + tl.arange(0, BSM)) % M
on = (pn * BSN + tl.arange(0, BSN)) % N
okp = tl.arange(0, BSK // 2)
okb = tl.arange(0, BSK)
oksb = pk * SPBS + okb
ap = a_ptr + om[:, None] * sa0 + oksb[None, :] * sa1
acc = tl.zeros((BSM, BSN), dtype=tl.float32)
for ki in range(pk * nki, (pk + 1) * nki):
inner = ki - pk * nki
oksp = pk * (SPBS // 2) + inner * (BSK // 2) + okp
oks = pk * (SPBS // SGS) + inner * (BSK // SGS) + tl.arange(0, BSK // SGS)
n_l = on[None, :].to(tl.int64)
k_l = oksp[:, None].to(tl.int64)
phys_b = _b_shuffle_phys_flat(n_l, k_l, KP)
bn = phys_b // KP
bk = phys_b % KP
bp = b_ptr + bn * sb0 + bk * sb1
ml = on[:, None].to(tl.int64)
nl = oks[None, :].to(tl.int64)
phys_s = _scale_shuffle_phys_flat(ml, nl, SN)
sr = phys_s // SN
sc = phys_s % SN
bsp = bs_ptr + sr * sbs0 + sc * sbs1
mask_sc = oks[None, :] < NUM_SG
if EVEN_K:
ab = tl.load(ap).to(tl.float32)
else:
ab = tl.load(ap, mask=okb[None, :] < (KP * 2) - ki * BSK, other=0.0).to(tl.float32)
af, asc = _mxfp4_quant_inline(ab, BSK, BSM, SGS)
bsc = tl.load(bsp, mask=mask_sc, other=127)
if EVEN_K:
bv = tl.load(bp)
else:
bv = tl.load(bp, mask=okp[:, None] < KP - ki * (BSK // 2), other=0)
acc = tl.dot_scaled(af, asc, "e2m1", bv, bsc, "e2m1", acc)
ap += BSK * sa1
c = acc.to(c_ptr.type.element_ty)
ocm = pm * BSM + tl.arange(0, BSM).to(tl.int64)
ocn = pn * BSN + tl.arange(0, BSN).to(tl.int64)
cp = c_ptr + scm * ocm[:, None] + scn * ocn[None, :] + pk * sck
cm = (ocm[:, None] < M) & (ocn[None, :] < N)
tl.store(cp, c, mask=cm)
@triton.jit
def _red_kernel(
yp, yo, M, N,
syk, sym, syn, som, son,
BM: tl.constexpr, BN: tl.constexpr,
NKS: tl.constexpr, NKP: 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)
mk = (om[:, None] < M) & (on[None, :] < N)
a = tl.zeros((BM, BN), dtype=tl.float32)
for k in range(NKS):
v = tl.load(yp + k * syk + om[:, None] * sym + on[None, :] * syn, mask=mk, other=0.0)
a += v
tl.store(yo + om[:, None] * som + on[None, :] * son, a.to(yo.type.element_ty), mask=mk)
def _get_spk(M, N, KP, BSK):
CU = 304
bm = 16 if M <= 16 else (32 if M <= 32 else (64 if M <= 64 else 128))
t = ((M + bm - 1) // bm) * ((N + 127) // 128)
c = CU / max(t, 1)
s = 0
while c >= pow(2, s + 1) and (pow(2, s + 1) * BSK) < 2 * KP:
s += 1
return min(s, 3)
def _spk_bs(KP, BSK, NKS):
if NKS <= 1:
return 2 * KP, BSK, 1
SP = triton.cdiv((2 * triton.cdiv(KP, NKS)), BSK) * BSK
b, n = BSK, NKS
while n > 1 and b > 16:
if KP % (SP // 2) == 0 and SP % b == 0 and KP % (b // 2) == 0:
break
elif KP % (SP // 2) != 0 and n > 1:
n //= 2
elif SP % b != 0:
if n > 1:
n //= 2
elif b > 16:
b //= 2
elif KP % (b // 2) != 0 and b > 16:
b //= 2
else:
break
SP = triton.cdiv((2 * triton.cdiv(KP, n)), b) * b
n = triton.cdiv(KP, (SP // 2))
return SP, b, n
def _run_separate_path(A, B_q, B_scale_sh, B_shuffle, m, n, k):
"""For large M, use separate quant + aiter ASM GEMM (faster than fused for compute-bound cases)."""
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
A_fp4, A_scale = dynamic_mxfp4_quant(A)
A_scale_sh = e8m0_shuffle(A_scale)
A_q = A_fp4.view(dtypes.fp4x2)
A_sc = A_scale_sh.view(dtypes.fp8_e8m0)
return aiter.gemm_a4w4(A_q, B_shuffle, A_sc, B_scale_sh, dtype=dtypes.bf16, bpreshuffle=True)
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
m, k = A.shape
n = B_q.shape[0]
KP = k // 2
if m >= 128:
return _run_separate_path(A, B_q, B_scale_sh, B_shuffle, m, n, k)
B_u8 = B_shuffle.view(torch.uint8)
BS_u8 = B_scale_sh.view(torch.uint8)
sb0, sb1 = B_u8.stride()
sbs0, sbs1 = BS_u8.stride()
sn = B_scale_sh.shape[1]
num_sg = k // SCALE_GROUP_SIZE
if m <= 16:
BSM, BSN, BSK = 16, 128, 256
GSM, nw, ns, wpe, mid = 1, 4, 2, 3, 16
NKS = 2 ** _get_spk(m, n, KP, BSK)
elif m <= 32:
BSM, BSN, BSK = 32, 128, 256
GSM, nw, ns, wpe, mid = 1, 4, 2, 3, 16
NKS = 1
elif m <= 64:
BSM, BSN, BSK = 64, 256, 256
GSM, nw, ns, wpe, mid = 1, 4, 3, 2, 32
NKS = 1
else:
BSM, BSN, BSK = 128, 256, 256
GSM, nw, ns, wpe, mid = 2, 4, 3, 2, 32
NKS = 1
BSK = max(BSK, 128)
if BSK >= 2 * KP:
BSK = triton.next_power_of_2(2 * KP)
NKS = 1
if NKS > 1:
SPBS, BSK, NKS = _spk_bs(KP, BSK, NKS)
else:
SPBS = 2 * KP
if NKS > 1:
ypp = torch.empty((NKS, m, n), dtype=torch.float32, device=A.device)
out = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
tgt = ypp
else:
out = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
ypp = None
tgt = out
grid = lambda META: (
META['NKS'] * triton.cdiv(m, META['BSM']) * triton.cdiv(n, META['BSN']),
)
_fqg_kernel[grid](
A, B_u8, tgt, BS_u8,
m, n, KP, sn, num_sg,
A.stride(0), A.stride(1), sb0, sb1, sbs0, sbs1,
0 if NKS == 1 else ypp.stride(0),
tgt.stride(-2) if NKS <= 1 else ypp.stride(1),
tgt.stride(-1) if NKS <= 1 else ypp.stride(2),
BSM=BSM, BSN=BSN, BSK=BSK,
GSM=GSM, NKS=NKS, SPBS=SPBS,
num_warps=nw, num_stages=ns, waves_per_eu=wpe, matrix_instr_nonkdim=mid,
)
if NKS > 1:
RBM, RBN = 16, 64
_red_kernel[(triton.cdiv(m, RBM), triton.cdiv(n, RBN))](
ypp, out, m, n,
ypp.stride(0), ypp.stride(1), ypp.stride(2),
out.stride(0), out.stride(1),
BM=RBM, BN=RBN, NKS=NKS, NKP=triton.next_power_of_2(NKS),
)
return out
scrolls · 348 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