submission 720499
Lemonade · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 239 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-720499?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:2238c1fff5e923cea2085a91f91abd99c7c1222e18cc51b64e578a0745988a04
license declaredunknown
license concludedunknown
authorsLemonade
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MXFP4-MM v27: Maximum performance with pre-allocated buffers and minimal overhead.split-k
GROUP_SIZE_M: tl.constexpr, NUM_KSPLIT: tl.constexpr, SPLITK_SIZE: tl.constexpr,stages = 2
NUM_KSPLIT=1, SPLITK_SIZE=k, num_warps=nw, num_stages=2)tile-k = 512
- Fused quant+GEMM for m<=32 with BK=512, BN=128 for m<=16 large Ktile-m = 16
BM = 16; nw = 2tile-n = 128
- Fused quant+GEMM for m<=32 with BK=512, BN=128 for m<=16 large KKernel source
submission.py239 lines
"""
MXFP4-MM v27: Maximum performance with pre-allocated buffers and minimal overhead.
- Pre-allocates ALL intermediate buffers on first call (zero alloc on hot path)
- Caches B reshapes across calls with same shape
- Fused quant+GEMM for m<=32 with BK=512, BN=128 for m<=16 large K
- CK GEMM for m>32 with pre-allocated output
- All view/reshape operations cached
"""
from task import input_t, output_t
import os
os.environ["CU_NUM"] = "256"
import torch
import triton
import triton.language as tl
# Global buffer cache - avoids torch.empty() on hot path
_cache = {}
@triton.jit
def _fused_qgemm(
a_ptr, b_ptr, c_ptr, b_scale_ptr,
M, N, K,
stride_am, stride_ak, stride_bn, stride_bk,
stride_ck, stride_cm, stride_cn, stride_bsn, stride_bsk,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr, NUM_KSPLIT: tl.constexpr, SPLITK_SIZE: tl.constexpr,
NUM_XCDS: tl.constexpr = 8,
):
SCALE_GROUP_SIZE: tl.constexpr = 32
tl.assume(stride_am > 0); tl.assume(stride_ak > 0)
tl.assume(stride_bn > 0); tl.assume(stride_bk > 0)
tl.assume(stride_cm > 0); tl.assume(stride_cn > 0)
tl.assume(stride_bsn > 0); tl.assume(stride_bsk > 0)
pid_raw = tl.program_id(0)
num_pid_m = tl.cdiv(M, BLOCK_M); num_pid_n = tl.cdiv(N, BLOCK_N)
GRID_MN = num_pid_m * num_pid_n; GT = GRID_MN * NUM_KSPLIT
ppx = (GT + NUM_XCDS - 1) // NUM_XCDS
tx = GT % NUM_XCDS; tx = NUM_XCDS if tx == 0 else tx
xcd = pid_raw % NUM_XCDS; lp = pid_raw // NUM_XCDS
if xcd < tx: pu = xcd * ppx + lp
else: pu = tx * ppx + (xcd - tx) * (ppx - 1) + lp
pid_k = pu % NUM_KSPLIT; pid = pu // NUM_KSPLIT
if NUM_KSPLIT == 1 and GROUP_SIZE_M > 1:
npig = GROUP_SIZE_M * num_pid_n; gid = pid // npig
fpm = gid * GROUP_SIZE_M; gsm = min(num_pid_m - fpm, GROUP_SIZE_M)
tl.assume(gsm >= 0)
pid_m = fpm + (pid % gsm); pid_n = (pid % npig) // gsm
else:
pid_m = pid // num_pid_n; pid_n = pid % num_pid_n
tl.assume(pid_m >= 0); tl.assume(pid_n >= 0); tl.assume(pid_k >= 0)
if (pid_k * SPLITK_SIZE // 2) < K:
nki = tl.cdiv(SPLITK_SIZE // 2, BLOCK_K // 2)
offs_m = (pid_m * BLOCK_M + tl.arange(0, BLOCK_M)) % M
ok_bf = pid_k * SPLITK_SIZE + tl.arange(0, BLOCK_K)
a_ptrs = a_ptr + offs_m[:, None] * stride_am + ok_bf[None, :] * stride_ak
obn = (pid_n * (BLOCK_N // 16) + tl.arange(0, BLOCK_N // 16)) % N
oks = pid_k * (SPLITK_SIZE // 2) * 16 + tl.arange(0, (BLOCK_K // 2) * 16)
b_ptrs = b_ptr + obn[:, None] * stride_bn + oks[None, :] * stride_bk
obsn = (pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)) % N
obsk = pid_k * (SPLITK_SIZE // SCALE_GROUP_SIZE * 32) + tl.arange(0, BLOCK_K // SCALE_GROUP_SIZE * 32)
bs_ptrs = b_scale_ptr + obsn[:, None] * stride_bsn + obsk[None, :] * stride_bsk
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
DM: tl.constexpr = 149 << 23
DF: tl.constexpr = tl.cast(DM, tl.float32, bitcast=True)
for _ in range(pid_k * nki, (pid_k + 1) * nki):
ab = tl.load(a_ptrs)
af = ab.reshape(BLOCK_M * (BLOCK_K // 32), 32).to(tl.float32)
ax = tl.max(tl.abs(af), axis=1, keep_dims=True)
ax = ax.to(tl.int32, bitcast=True)
ax = (ax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
ax = ax.to(tl.float32, bitcast=True)
su = tl.log2(ax).floor() - 2
su = tl.clamp(su, min=-127, max=127)
a_sc = (su.to(tl.uint8) + 127).reshape(BLOCK_M, BLOCK_K // 32)
qx = af * tl.exp2(-su)
qx = qx.to(tl.uint32, bitcast=True); sg = qx & 0x80000000; qx = qx ^ sg
qf = qx.to(tl.float32, bitcast=True)
st = qf >= 6; dn = (not st) & (qf < 1); nr = not (st | dn)
dx = (qf + DF).to(tl.uint32, bitcast=True) - DM; dx = dx.to(tl.uint8)
nx = qx.to(tl.int32, bitcast=True); mo = (nx >> 22) & 1
nx += (-126 << 23) + (1 << 21) - 1; nx += mo; nx = (nx >> 22).to(tl.uint8)
e = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)
e = tl.where(nr, nx, e); e = tl.where(dn, dx, e)
e = e | (sg >> 28).to(tl.uint8)
e = tl.reshape(e, [BLOCK_M * (BLOCK_K // 32), 16, 2])
ev, od = tl.split(e); afp4 = (ev | (od << 4)).reshape(BLOCK_M, BLOCK_K // 2)
br = tl.load(b_ptrs)
b = (br.reshape(1, BLOCK_N // 16, BLOCK_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5).reshape(BLOCK_N, BLOCK_K // 2).trans(1, 0))
bsc = (tl.load(bs_ptrs)
.reshape(BLOCK_N // 32, BLOCK_K // 32 // 8, 4, 16, 2, 2, 1)
.permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_N, BLOCK_K // 32))
acc = tl.dot_scaled(afp4, a_sc, "e2m1", b, bsc, "e2m1", acc)
a_ptrs += BLOCK_K * stride_ak
b_ptrs += (BLOCK_K // 2) * 16 * stride_bk
bs_ptrs += BLOCK_K * stride_bsk
ocm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M).to(tl.int64)
ocn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N).to(tl.int64)
cm = (ocm[:, None] < M) & (ocn[None, :] < N)
if NUM_KSPLIT == 1:
tl.store(c_ptr + ocm[:, None] * stride_cm + ocn[None, :] * stride_cn,
acc.to(c_ptr.type.element_ty), mask=cm)
else:
tl.store(c_ptr + pid_k * stride_ck + ocm[:, None] * stride_cm + ocn[None, :] * stride_cn,
acc, mask=cm)
@triton.jit
def _reduce(cp, co, M, N, spk, spm, spn, som, son,
BM: tl.constexpr, BN: tl.constexpr, NK: tl.constexpr, MK: tl.constexpr):
pm = tl.program_id(0); pn = tl.program_id(1)
om = (pm * BM + tl.arange(0, BM)) % M; on = (pn * BN + tl.arange(0, BN)) % N
ok = tl.arange(0, MK)
p = cp + ok[:, None, None] * spk + om[None, :, None] * spm + on[None, None, :] * spn
v = tl.load(p, mask=ok[:, None, None] < NK) if NK != MK else tl.load(p)
tl.store(co + om[:, None] * som + on[None, :] * son, tl.sum(v, axis=0).to(co.type.element_ty))
@triton.jit
def _qshuf(x_ptr, xf_ptr, bs_ptr, sxm, sxn, sfm, sfn, sbm, sbn,
M: tl.constexpr, N: tl.constexpr, sN: tl.constexpr,
sMP: tl.constexpr, sNP: tl.constexpr,
BS: tl.constexpr, QBS: tl.constexpr):
pm = tl.program_id(0); pn = tl.program_id(1)
sxm = tl.cast(sxm, tl.int64); sxn = tl.cast(sxn, tl.int64)
sfm = tl.cast(sfm, tl.int64); sfn = tl.cast(sfn, tl.int64)
xm = pm * BS + tl.arange(0, BS); xn = pn * QBS + tl.arange(0, QBS)
x = tl.load(x_ptr + xm[:, None] * sxm + xn[None, :] * sxn,
mask=(xm < M)[:, None] & (xn < N)[None, :]).to(tl.float32)
ax = tl.max(tl.abs(x), axis=1, keep_dims=True)
ax = ax.to(tl.int32, bitcast=True)
ax = (ax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
ax = ax.to(tl.float32, bitcast=True)
su = tl.log2(ax).floor() - 2; su = tl.clamp(su, min=-127, max=127)
qx = x * tl.exp2(-su); bs = su.to(tl.uint8) + 127
qx = qx.to(tl.uint32, bitcast=True); s = qx & 0x80000000; qx = qx ^ s
qf = qx.to(tl.float32, bitcast=True)
st = qf >= 6; dn = (not st) & (qf < 1); nr = not (st | dn)
DM: tl.constexpr = 149 << 23; DF: tl.constexpr = tl.cast(DM, tl.float32, bitcast=True)
dx = (qf + DF).to(tl.uint32, bitcast=True) - DM; dx = dx.to(tl.uint8)
nx = qx.to(tl.int32, bitcast=True); mo = (nx >> 22) & 1
nx += (-126 << 23) + (1 << 21) - 1; nx += mo; nx = (nx >> 22).to(tl.uint8)
e = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)
e = tl.where(nr, nx, e); e = tl.where(dn, dx, e); e = e | (s >> 28).to(tl.uint8)
e = tl.reshape(e, [BS, QBS // 2, 2]); ev, od = tl.split(e); ot = ev | (od << 4)
om = pm * BS + tl.arange(0, BS); on = pn * QBS // 2 + tl.arange(0, QBS // 2)
tl.store(xf_ptr + om[:, None] * sfm + on[None, :] * sfn, ot,
mask=(om < M)[:, None] & (on < (N // 2))[None, :])
bm = pm * BS + tl.arange(0, BS); bn = pn
b0 = bm[:, None] // 32; b12 = bm[:, None] % 32; b1 = b12 // 16; b2 = b12 % 16
b3 = bn[None, :] // 8; b45 = bn[None, :] % 8; b4 = b45 // 4; b5 = b45 % 4
bo = b1 + b4*2 + b2*4 + b5*64 + b3*256 + b0*32*sN
m1 = (bm < M)[:, None] & (bn < sN)[None, :]
m2 = (bm < sMP)[:, None] & (bn < sNP)[None, :]
tl.store(bs_ptr + bo, tl.where(m1, bs, 127), mask=m2)
def _get_bufs(m, n, k, device):
"""Get pre-allocated buffers for a given shape. Avoids torch.empty() on hot path."""
key = (m, n, k)
if key not in _cache:
from aiter import dtypes
sM = triton.cdiv(m, 32) * 32
sNv = triton.cdiv(k, 32)
sN = triton.cdiv(sNv, 8) * 8
_cache[key] = {
'xf': torch.empty((m, k // 2), dtype=torch.uint8, device=device),
'bs': torch.empty((triton.cdiv(m, 256) * 256, sN), dtype=torch.uint8, device=device),
'C': torch.empty((m, n), dtype=torch.bfloat16, device=device),
'sNv': sNv, 'sM': sM, 'sN': sN,
}
return _cache[key]
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
# Cache B reshapes (same across calls)
bkey = (id(B_shuffle), id(B_scale_sh))
if bkey not in _cache:
_cache[bkey] = (
B_shuffle.view(torch.uint8).reshape(n // 16, (k // 2) * 16),
B_scale_sh.view(torch.uint8).reshape(B_scale_sh.view(torch.uint8).shape[0] // 32,
B_scale_sh.view(torch.uint8).shape[1] * 32),
)
Bp, Bs = _cache[bkey]
BN = 64
if m <= 32:
BM = 16; nw = 2
BK = 512 if k >= 512 and k % 512 == 0 else 256
if BK == 512: nw = 4
if m <= 16 and k > 1024: BN = 128
tiles = triton.cdiv(m, BM) * triton.cdiv(n, BN)
NS = 1; mxk = k // (BK * 2)
if mxk >= 2 and tiles < 128:
while tiles * NS <= 256 and NS < mxk: NS *= 2
NS = min(NS, max(1, min(mxk, 8)))
SS = triton.cdiv(triton.cdiv(k, NS), BK) * BK
AK = triton.cdiv(k, SS)
grid = (AK * triton.cdiv(m, BM) * triton.cdiv(n, BN),)
if AK == 1:
C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
_fused_qgemm[grid](A, Bp, C, Bs, m, n, k,
A.stride(0), A.stride(1), Bp.stride(0), Bp.stride(1),
0, C.stride(0), C.stride(1), Bs.stride(0), Bs.stride(1),
BLOCK_M=BM, BLOCK_N=BN, BLOCK_K=BK, GROUP_SIZE_M=4,
NUM_KSPLIT=1, SPLITK_SIZE=k, num_warps=nw, num_stages=2)
return C
else:
Cp = torch.empty((AK, m, n), dtype=torch.float32, device=A.device)
_fused_qgemm[grid](A, Bp, Cp, Bs, m, n, k,
A.stride(0), A.stride(1), Bp.stride(0), Bp.stride(1),
Cp.stride(0), Cp.stride(1), Cp.stride(2), Bs.stride(0), Bs.stride(1),
BLOCK_M=BM, BLOCK_N=BN, BLOCK_K=BK, GROUP_SIZE_M=1,
NUM_KSPLIT=AK, SPLITK_SIZE=SS, num_warps=nw, num_stages=2)
C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
MK = triton.next_power_of_2(AK)
_reduce[(triton.cdiv(m, 16), triton.cdiv(n, 64))](
Cp, C, m, n, Cp.stride(0), Cp.stride(1), Cp.stride(2),
C.stride(0), C.stride(1), BM=16, BN=64, NK=AK, MK=MK)
return C
else:
# Fast quant+shuffle + CK GEMM with pre-allocated buffers
from aiter import dtypes
import aiter
bufs = _get_bufs(m, n, k, A.device)
xf, bs_buf = bufs['xf'], bufs['bs']
sNv, sM, sN = bufs['sNv'], bufs['sM'], bufs['sN']
_qshuf[(triton.cdiv(m, 128), sNv)](
A, xf, bs_buf, A.stride(0), A.stride(1), xf.stride(0), xf.stride(1),
bs_buf.stride(0), bs_buf.stride(1), M=m, N=k, sN=sNv, sMP=sM, sNP=sN, BS=128, QBS=32)
return aiter.gemm_a4w4(xf.view(dtypes.fp4x2), B_shuffle,
bs_buf.view(dtypes.fp8_e8m0), B_scale_sh,
dtype=dtypes.bf16, bpreshuffle=True)
scrolls · 239 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