submission 753771
Aniket Sadashiva · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 237 lines, June 9 Researcher Reciprocity License v1.0.
submission_vh366.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-753771?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:0d80c6ce6866c25a1ddda3824b0a9ac0b112305486157e334dc1300347834d20
license declaredunknown
license concludedunknown
authorsAniket Sadashiva
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 4
vh366: vh362 base + quant kernel tuning (num_warps=4 instead of 2).tile-k = 256
BM=16, BN=64, BK=256, GM=1, NKS=1, SBS=kp<<1, EK=True,tile-m = 16
BM=16, BN=64, BK=256, GM=1, NKS=1, SBS=kp<<1, EK=True,tile-n = 64
BM=16, BN=64, BK=256, GM=1, NKS=1, SBS=kp<<1, EK=True,Kernel source
submission_vh366.py237 lines
"""
vh366: vh362 base + quant kernel tuning (num_warps=4 instead of 2).
Also try BS=32 for the quant kernel (more blocks for M=64).
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
import aiter
from aiter import dtypes
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
try:
import triton._utils as _tu
_d = _tu.type_canonicalisation_dict
_d.setdefault("float4_e2m1fn_x2", "u8")
_d.setdefault("float8_e8m0fnu", "u8")
_d.setdefault("float4_e2m1fn", "u8")
except Exception:
pass
try:
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
HAS_ASM = True
except ImportError:
gemm_a4w4_asm = None
HAS_ASM = False
_KN = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
@triton.jit
def _pg(pid: int, npm: int, npn: int, GM: tl.constexpr = 1):
if GM == 1: return pid // npn, pid % npn
nig = GM * npn; gid = pid // nig; fpm = gid * GM
gsm = min(npm - fpm, GM); tl.assume(gsm >= 0)
return fpm + (pid % gsm), (pid % nig) // gsm
@triton.jit
def _gk(ap, bp, cp, bsp, M, N, K,
sam, sak, sbk, sbn, sck, scm, scn, sbsk, sbsn,
BM: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr, GM: tl.constexpr,
NKS: tl.constexpr, SBS: tl.constexpr, EK: tl.constexpr,
PRESHUFFLE: tl.constexpr):
tl.assume(sam > 0); tl.assume(sak > 0); tl.assume(sbk > 0); tl.assume(sbn > 0)
tl.assume(scm > 0); tl.assume(scn > 0); tl.assume(sbsk > 0); tl.assume(sbsn > 0)
pu = tl.program_id(0); pk = pu % NKS; p = pu // NKS
npm = tl.cdiv(M, BM); npn = tl.cdiv(N, BN)
if NKS == 1: pm, pn = _pg(p, npm, npn, GM=GM)
else: pm = p // npn; pn = p % npn
tl.assume(pm >= 0); tl.assume(pn >= 0); tl.assume(pk >= 0)
SG: tl.constexpr = 32; ST: tl.constexpr = BK // SG
if (pk * SBS // 2) < K:
nki = tl.cdiv(SBS // 2, BK // 2)
okb = tl.arange(0, BK); oksb = pk * SBS + okb
oam = (pm * BM + tl.arange(0, BM)) % M
apt = ap + (oam[:, None] * sam + oksb[None, :] * sak)
if PRESHUFFLE:
obn_ps = (pn * (BN // 16) + tl.arange(0, BN // 16)) % (N // 16)
oks_ps = pk * (SBS // 2) * 16 + tl.arange(0, (BK // 2) * 16)
bpt = bp + obn_ps[:, None] * sbn + oks_ps[None, :] * sbk
obsn = (pn * (BN // 32) + tl.arange(0, BN // 32)) % (N // 32)
obsk = (pk * (SBS // SG) * 32) + tl.arange(0, BK // SG * 32)
bspt = bsp + obsn[:, None] * sbsn + obsk[None, :] * sbsk
else:
ok = tl.arange(0, BK // 2); oks_nat = pk * (SBS // 2) + ok
obn_nat = (pn * BN + tl.arange(0, BN)) % N
bpt = bp + (obn_nat[:, None] * sbn + oks_nat[None, :] * sbk)
ok2 = pk * (SBS // SG) + tl.arange(0, ST)
d0 = obn_nat // 32; d1 = (obn_nat & 31) >> 4; d2 = obn_nat & 15
srp = d0 * (32 * sbsn) + d2 * 4 + d1
d3 = ok2 >> 3; d4 = (ok2 & 7) >> 2; d5 = ok2 & 3
scp2 = d3 * 256 + d5 * 64 + d4 * 2
bso = srp[:, None] + scp2[None, :]
acc = tl.zeros((BM, BN), dtype=tl.float32)
for ki in range(0, nki):
if PRESHUFFLE:
bs = (tl.load(bspt)
.reshape(BN // 32, BK // SG // 8, 4, 16, 2, 2, 1)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BN, BK // SG))
if EK: ab = tl.load(apt); b_raw = tl.load(bpt)
else:
ab = tl.load(apt, mask=okb[None, :] < SBS, other=0)
b_raw = tl.load(bpt, mask=(obn_ps[:, None] < (N // 16)) & (oks_ps[None, :] < (K * 16)), other=0)
b = (b_raw.reshape(1, BN // 16, BK // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BN, BK // 2)
.trans(1, 0))
else:
aki = pk * nki + ki
bs = tl.load(bsp + bso)
if EK: ab = tl.load(apt); b = tl.load(bpt).trans(1, 0)
else:
ab = tl.load(apt, mask=okb[None, :] < 2 * K - aki * BK, other=0)
b = tl.load(bpt, mask=tl.arange(0, BK // 2)[None, :] < K - aki * (BK // 2), other=0).trans(1, 0)
a, asc = _mxfp4_quant_op(ab, BK, BM, 32)
acc += tl.dot_scaled(a, asc, "e2m1", b, bs, "e2m1")
apt += BK * sak
if PRESHUFFLE:
bpt += (BK // 2) * 16 * sbk; bspt += BK * sbsk
else:
bpt += (BK // 2) * sbk; ok2 += ST
d3 = ok2 >> 3; d4 = (ok2 & 7) >> 2; d5 = ok2 & 3
scp2 = d3 * 256 + d5 * 64 + d4 * 2
bso = srp[:, None] + scp2[None, :]
c = acc.to(cp.type.element_ty)
ocm = pm * BM + tl.arange(0, BM).to(tl.int64)
ocn = pn * BN + tl.arange(0, BN).to(tl.int64)
cpt = cp + scm * ocm[:, None] + scn * ocn[None, :] + pk * sck
cm = (ocm[:, None] < M) & (ocn[None, :] < N)
if NKS == 1: tl.store(cpt, c, mask=cm, cache_modifier=".wt")
else: tl.store(cpt, c, mask=cm)
@triton.jit
def _rk(src, dst, M, N, ss, sm, sn, dm, dn,
NKS: tl.constexpr, BM: tl.constexpr, BN: tl.constexpr):
pid = tl.program_id(0); npn = tl.cdiv(N, BN); pm = pid // npn; pn = pid % npn
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 s in tl.static_range(NKS):
a += tl.load(src + s * ss + om[:, None] * sm + on[None, :] * sn, mask=mk, other=0.0)
tl.store(dst + om[:, None] * dm + on[None, :] * dn, a.to(tl.bfloat16), mask=mk, cache_modifier=".wt")
@triton.jit
def _qk(xp, fp, bp, sxm, sxn, sfm, sfn, M, N, scN, sMp, sNp, BS: tl.constexpr):
qb: tl.constexpr = 32; pm = tl.program_id(0); pn = tl.program_id(1)
sxm64 = tl.cast(sxm, tl.int64); sxn64 = tl.cast(sxn, tl.int64)
sfm64 = tl.cast(sfm, tl.int64); sfn64 = tl.cast(sfn, tl.int64)
xom = pm * BS + tl.arange(0, BS); xon = pn * qb + tl.arange(0, qb)
xmk = (xom < M)[:, None] & (xon < N)[None, :]
x = tl.load(xp + xom[:, None] * sxm64 + xon[None, :] * sxn64, mask=xmk).to(tl.float32)
xf, be = _mxfp4_quant_op(x, qb, BS, qb)
oon = pn * (qb // 2) + tl.arange(0, qb // 2)
omk = (xom < M)[:, None] & (oon < (N // 2))[None, :]
tl.store(fp + xom[:, None] * sfm64 + oon[None, :] * sfn64, xf, mask=omk)
bv = tl.reshape(be, [BS]); bm = xom; bn = pn
d0 = bm // 32; r32 = bm % 32; d2 = r32 % 16; d1 = r32 // 16
d3 = bn // 8; r8 = bn % 8; d5 = r8 % 4; d4 = r8 // 4
so = d1 + d4 * 2 + d2 * 4 + d5 * 64 + d3 * 256 + d0 * (32 * scN)
m1 = (bm < M) & (bn < scN); m2 = (bm < sMp) & (bn < sNp)
bv = tl.where(m1, bv, 127); tl.store(bp + so, bv, mask=m2)
_C = {}; _W = set()
def _dk(d): return d.type, d.index
def _ct(n, s, dt, d):
k = (n, _dk(d)); o = _C.get(k)
if o is None: o = torch.empty(s, dtype=dt, device=d); _C[k] = o
return o
def _ab(m, n, k, d):
key = (("a", m, n, k), _dk(d)); c = _C.get(key)
if c: return c
sv = triton.cdiv(k, 32); sp = triton.cdiv(sv, 8) * 8; sm = triton.cdiv(m, 32) * 32
xf = torch.empty((m, k // 2), dtype=torch.uint8, device=d)
bs = torch.empty((sm, sp), dtype=torch.uint8, device=d)
out = torch.empty((((m + 31) // 32) * 32, n), dtype=torch.bfloat16, device=d)
xf_fp4 = xf.view(dtypes.fp4x2).view(m, k // 2); bs_e8m0 = bs.view(dtypes.fp8_e8m0)
c = (xf, bs, xf_fp4, bs_e8m0, out, sv, sm, sp); _C[key] = c; return c
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]
kp = k >> 1
if k == 512:
b_q_u8 = b_q.view(torch.uint8)
b_scale_u8 = b_scale_sh.view(torch.uint8)
out = _ct(("o", m, n, k), (m, n), torch.bfloat16, a.device)
g = triton.cdiv(m, 16) * triton.cdiv(n, 64)
_gk[(g,)](a, b_q_u8, out, b_scale_u8, m, n, kp,
a.stride(0), a.stride(1), b_q_u8.stride(1), b_q_u8.stride(0),
0, out.stride(0), out.stride(1), 1, b_scale_u8.shape[1],
BM=16, BN=64, BK=256, GM=1, NKS=1, SBS=kp<<1, EK=True,
PRESHUFFLE=False, num_warps=4)
return out
if m == 16 and k == 7168:
bsu8 = b_shuffle.view(torch.uint8)
bt = bsu8.view(n // 16, kp * 16)
bscu8 = b_scale_sh.view(torch.uint8)
bst = bscu8.reshape(bscu8.shape[0] // 32, bscu8.shape[1] * 32)
NS = 7; SBS = (2 * kp) // NS
skb = _ct(("sk", m, n, k), (NS, m, n), torch.float32, a.device)
out = _ct(("so", m, n, k), (m, n), torch.bfloat16, a.device)
g = triton.cdiv(n, 128)
_gk[(g * NS,)](a, bt, skb, bst, m, n, kp,
a.stride(0), a.stride(1), bt.stride(1), bt.stride(0),
m * n, skb.stride(1), skb.stride(2), bst.stride(1), bst.stride(0),
BM=16, BN=128, BK=512, GM=1, NKS=NS, SBS=SBS, EK=True,
PRESHUFFLE=True, num_warps=4)
_rk[(g,)](skb, out, m, n, skb.stride(0), skb.stride(1), skb.stride(2),
out.stride(0), out.stride(1), NKS=NS, BM=16, BN=128, num_warps=4)
return out
# M=64/M=256: quant with num_warps=4 (was 2) + ASM
if HAS_ASM:
xf, bs, xf_fp4, bs_e8m0, out, sv, sm, sp = _ab(m, n, k, a.device)
_qk[(((m + 63) >> 6), sv)](a, xf, bs,
a.stride(0), a.stride(1), xf.stride(0), xf.stride(1),
M=m, N=k, scN=sv, sMp=sm, sNp=sp, BS=64, num_warps=4)
gemm_a4w4_asm(xf_fp4, b_shuffle, bs_e8m0, b_scale_sh, out, _KN,
bpreshuffle=True, log2_k_split=1)
return out[:m]
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
afp4, asc = dynamic_mxfp4_quant(a)
return aiter.gemm_a4w4(afp4.view(dtypes.fp4x2), b_shuffle,
e8m0_shuffle(asc).view(dtypes.fp8_e8m0), b_scale_sh, dtype=dtypes.bf16, bpreshuffle=True)
def _ms(m, n, k, dev):
from aiter.ops.shuffle import shuffle_weight
a = torch.zeros((m, k), dtype=torch.bfloat16, device=dev)
b = torch.zeros((n, k), dtype=torch.bfloat16, device=dev)
bq = torch.zeros((n, k // 2), dtype=torch.uint8, device=dev).view(dtypes.fp4x2)
sv = triton.cdiv(k, 32); sp = triton.cdiv(sv, 8) * 8; sm = triton.cdiv(n, 32) * 32
bss = torch.full((sm, sp), 127, dtype=torch.uint8, device=dev).view(dtypes.fp8_e8m0)
return a, b, bq, shuffle_weight(bq, layout=(16, 16)), bss
def _pw():
if not torch.cuda.is_available(): return
dev = torch.device("cuda"); dk = _dk(dev)
if dk in _W: return
_W.add(dk)
for m, n, k in [(4,2880,512),(32,4096,512),(32,2880,512),(16,2112,7168),(64,7168,2048),(256,3072,1536)]:
try: custom_kernel(_ms(m, n, k, dev))
except Exception as e: print(f"[vh366] prewarm ({m},{n},{k}): {e}", flush=True)
try: torch.cuda.synchronize(dev)
except: pass
try: _pw()
except: pass
scrolls · 237 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