submission 754367
flower2123 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 414 lines, June 9 Researcher Reciprocity License v1.0.
submission_v1_flower_mm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-754367?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:2e655cfed189ae7fcb9826895c4d3ccc40e1f198360c6ff601aeba04a60a0990
license declaredunknown
license concludedunknown
authorsflower2123
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
submission_v1_flower_mm.py414 lines
import torch
import triton
import triton.language as tl
from task import input_t, output_t
_out_pool = {}
_params = {}
_grid_memo = {}
PER_SHAPE = {
(4, 2880, 512): dict(tm=4, tn=128, tk=256, gm=1, warp=4, pipe=2, occ=1, kdim=16, cmod=None, ks=1),
(16, 2112, 7168): dict(tm=16, tn=128, tk=512, gm=1, warp=4, pipe=2, occ=3, kdim=16, cmod=".cg", ks=14),
(32, 4096, 512): dict(tm=16, tn=32, tk=256, gm=1, warp=4, pipe=3, occ=3, kdim=16, cmod=".cg", ks=1),
(32, 2880, 512): dict(tm=8, tn=128, tk=256, gm=1, warp=4, pipe=2, occ=2, kdim=16, cmod=None, ks=1),
(64, 7168, 2048): dict(tm=16, tn=128, tk=256, gm=1, warp=4, pipe=2, occ=2, kdim=16, cmod=".cg", ks=1),
(256, 3072, 1536): dict(tm=16, tn=256, tk=512, gm=1, warp=8, pipe=2, occ=2, kdim=16, cmod=None, ks=1),
}
FALLBACK = dict(tm=16, tn=32, tk=256, gm=1, warp=2, pipe=2, occ=0, kdim=16, cmod=".cg", ks=1)
SEP_SHAPE = {
(32, 4096, 512): dict(tm=32, tn=128, tk=256, gm=4, warp=4, pipe=2, occ=0, kdim=16, cmod=".cg", ks=1),
(32, 2880, 512): dict(tm=32, tn=64, tk=256, gm=4, warp=4, pipe=2, occ=0, kdim=16, cmod=".cg", ks=1),
(64, 7168, 2048): dict(tm=16, tn=128, tk=512, gm=4, warp=4, pipe=2, occ=0, kdim=16, cmod=".cg", ks=2),
}
SEP_FALLBACK = dict(tm=16, tn=64, tk=256, gm=4, warp=4, pipe=2, occ=0, kdim=16, cmod=".cg", ks=1)
def _adj_split(kp, bk, ns):
sb = triton.cdiv(2 * triton.cdiv(kp, ns), bk) * bk
while ns > 1 and bk > 16:
ok = kp % (sb // 2) == 0 and sb % bk == 0 and kp % (bk // 2) == 0
if ok:
break
if kp % (sb // 2) != 0 and ns > 1:
ns //= 2
elif sb % bk != 0:
ns = ns // 2 if ns > 1 else ns
if ns <= 1 and bk > 16:
bk //= 2
elif kp % (bk // 2) != 0 and bk > 16:
bk //= 2
else:
break
sb = triton.cdiv(2 * triton.cdiv(kp, ns), bk) * bk
return sb, bk, triton.cdiv(kp, sb // 2)
def _expand(raw, kp):
c = {
"BLOCK_SIZE_M": raw["tm"], "BLOCK_SIZE_N": max(raw["tn"], 32),
"BLOCK_SIZE_K": raw["tk"], "GROUP_SIZE_M": raw["gm"],
"num_warps": raw["warp"], "num_stages": raw["pipe"],
"waves_per_eu": raw["occ"], "matrix_instr_nonkdim": raw["kdim"],
"cache_modifier": raw["cmod"], "NUM_KSPLIT": raw["ks"],
}
if c["NUM_KSPLIT"] > 1:
sb, bk, ns = _adj_split(kp, c["BLOCK_SIZE_K"], c["NUM_KSPLIT"])
c.update(SPLITK_BLOCK_SIZE=sb, BLOCK_SIZE_K=bk, NUM_KSPLIT=ns)
else:
c.update(SPLITK_BLOCK_SIZE=2 * kp, NUM_KSPLIT=1)
if c["BLOCK_SIZE_K"] >= 2 * kp:
c.update(BLOCK_SIZE_K=triton.next_power_of_2(2 * kp), SPLITK_BLOCK_SIZE=2 * kp, NUM_KSPLIT=1)
return c
@triton.jit
def _fp4_encode(
src, R: tl.constexpr, C: tl.constexpr,
):
Q: tl.constexpr = 32
NB: tl.constexpr = C // Q
f32 = src.to(tl.float32).reshape(R, NB, Q)
mx = tl.max(tl.abs(f32), axis=-1, keep_dims=True)
mx = (mx.to(tl.int32, bitcast=True) + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
lx = ((mx >> 23) & 0xFF).to(tl.int32) - 127
ub = tl.minimum(tl.maximum(lx - 2, -127), 127)
e8 = ub.to(tl.uint8) + 127
hb = (ub.to(tl.int32) + 127).to(tl.uint32) << 23
hf = hb.to(tl.float32, bitcast=True)
hf_full = tl.broadcast_to(hf, (R, NB, Q)).reshape(R, C)
pv = hf_full.reshape(R, C // 2, 2)
ev, _ = tl.split(pv)
ev = ev.reshape(R, C // 2)
raw16 = src.to(tl.uint16, bitcast=True).reshape(R, C // 2, 2)
lo16, hi16 = tl.split(raw16)
p32 = lo16.to(tl.uint32) | (hi16.to(tl.uint32) << 16)
p32 = p32.reshape(R, C // 2)
enc = tl.inline_asm_elementwise(
"v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2", "=v, v, v",
[p32, ev], dtype=tl.uint32, is_pure=True, pack=1,
)
return (enc & 0xFF).to(tl.uint8).reshape(R, C // 2), e8.reshape(R, NB)
@triton.jit
def _encode_block(
inp, o_fp4, o_sc,
nrow, ncol,
si0, si1, sq0, sq1, ss0, ss1,
TR: tl.constexpr, TC: tl.constexpr,
):
r = tl.program_id(0) * TR + tl.arange(0, TR)
c = tl.program_id(1) * TC + tl.arange(0, TC)
dat = tl.load(inp + r[:, None] * si0 + c[None, :] * si1,
mask=(r[:, None] < nrow) & (c[None, :] < ncol), other=0.0)
f4, sc = _fp4_encode(dat, TR, TC)
HC: tl.constexpr = TC // 2
qc = tl.program_id(1) * HC + tl.arange(0, HC)
tl.store(o_fp4 + r[:, None] * sq0 + qc[None, :] * sq1,
f4, mask=(r[:, None] < nrow) & (qc[None, :] < ncol // 2))
SC: tl.constexpr = TC // 32
sc_c = tl.program_id(1) * SC + tl.arange(0, SC)
tl.store(o_sc + r[:, None] * ss0 + sc_c[None, :] * ss1,
sc, mask=(r[:, None] < nrow) & (sc_c[None, :] < ncol // 32))
@triton.heuristics({"OK": lambda a: (a["K"] % (a["BLOCK_SIZE_K"] // 2) == 0)
and (a["SPLITK_BLOCK_SIZE"] % a["BLOCK_SIZE_K"] == 0)
and (a["K"] % (a["SPLITK_BLOCK_SIZE"] // 2) == 0)})
@triton.jit
def _matmul_fused(
a_p, b_p, c_p, bs_p,
M, N, K,
sa0, sa1, sb0, sb1, sc_k, sc0, sc1, sbs0, sbs1,
BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr,
NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,
OK: tl.constexpr,
num_warps: tl.constexpr, num_stages: tl.constexpr,
waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
cache_modifier: tl.constexpr,
):
tl.assume(sa0 > 0); tl.assume(sa1 > 0)
tl.assume(sb0 > 0); tl.assume(sb1 > 0)
tl.assume(sc0 > 0); tl.assume(sc1 > 0)
tl.assume(sbs0 > 0); tl.assume(sbs1 > 0)
G: tl.constexpr = 32
nm = tl.cdiv(M, BLOCK_SIZE_M)
nn = tl.cdiv(N, BLOCK_SIZE_N)
uid = tl.program_id(0)
sk = uid % NUM_KSPLIT
flat = uid // NUM_KSPLIT
if NUM_KSPLIT == 1:
gn = GROUP_SIZE_M * nn
gi = flat // gn
fm = gi * GROUP_SIZE_M
gs = min(nm - fm, GROUP_SIZE_M)
im = fm + (flat % gn) % gs
jn = (flat % gn) // gs
else:
im = flat // nn
jn = flat % nn
tl.assume(im >= 0); tl.assume(jn >= 0); tl.assume(sk >= 0)
if (sk * SPLITK_BLOCK_SIZE // 2) < K:
niter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
rm = (im * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
ck = sk * SPLITK_BLOCK_SIZE + tl.arange(0, BLOCK_SIZE_K)
pa = a_p + rm[:, None] * sa0 + ck[None, :] * sa1
sh_a = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
sh_o = sk * (SPLITK_BLOCK_SIZE // 2) * 16 + sh_a
bn = (jn * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % (N // 16)
pb = b_p + bn[:, None] * sb0 + sh_o[None, :] * sb1
bsn = jn * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N // 32) * 32
bsk = sk * (SPLITK_BLOCK_SIZE // G) * 32 + tl.arange(0, BLOCK_SIZE_K // G * 32)
pbs = bs_p + bsn[:, None] * sbs0 + bsk[None, :] * sbs1
dot = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for ki in range(sk * niter, (sk + 1) * niter):
if OK:
va = tl.load(pa)
vbs = tl.load(pbs, cache_modifier=cache_modifier)
vb = tl.load(pb, cache_modifier=cache_modifier)
else:
lo = (ki - sk * niter) * BLOCK_SIZE_K
va = tl.load(pa, mask=tl.arange(0, BLOCK_SIZE_K)[None, :] < (2 * K - sk * SPLITK_BLOCK_SIZE - lo), other=0.0)
vbs = tl.load(pbs, cache_modifier=cache_modifier)
vb = tl.load(pb,
mask=sh_a[None, :] < ((K - (sk * (SPLITK_BLOCK_SIZE // 2) + (ki - sk * niter) * (BLOCK_SIZE_K // 2))) * 16),
other=0, cache_modifier=cache_modifier)
aq, asc = _fp4_encode(va, BLOCK_SIZE_M, BLOCK_SIZE_K)
ws = (vbs
.reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // G // 8, 4, 16, 2, 2, 1)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // G))
bd = (vb.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
.trans(1, 0))
dot = tl.dot_scaled(aq, asc, "e2m1", bd, ws, "e2m1", dot)
pa += BLOCK_SIZE_K * sa1
pb += (BLOCK_SIZE_K // 2) * 16 * sb1
pbs += BLOCK_SIZE_K * sbs1
out = dot.to(c_p.type.element_ty)
om = im * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
on = jn * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
tl.store(c_p + sc0 * om[:, None] + sc1 * on[None, :] + sk * sc_k,
out, mask=(om[:, None] < M) & (on[None, :] < N))
@triton.jit
def _merge(src, dst, M, N, sk0, sm0, sn0, dm0, dn0,
TM: tl.constexpr, TN: tl.constexpr,
REAL: tl.constexpr, CAP: tl.constexpr):
rm = (tl.program_id(0) * TM + tl.arange(0, TM)) % M
rn = (tl.program_id(1) * TN + tl.arange(0, TN)) % N
base = src + rm[:, None] * sm0 + rn[None, :] * sn0
s = tl.load(base).to(tl.float32)
for j in tl.static_range(1, CAP):
if j < REAL:
s += tl.load(base + j * sk0).to(tl.float32)
tl.store(dst + rm[:, None] * dm0 + rn[None, :] * dn0, s.to(dst.type.element_ty))
@triton.heuristics({"OK": lambda a: (a["K"] % (a["BLOCK_SIZE_K"] // 2) == 0)
and (a["SPLITK_BLOCK_SIZE"] % a["BLOCK_SIZE_K"] == 0)
and (a["K"] % (a["SPLITK_BLOCK_SIZE"] // 2) == 0)})
@triton.jit
def _matmul_sep(
q4, qsc, b_p, c_p, bs_p,
M, N, K,
sq0, sq1, ss0, ss1,
sb0, sb1, sc_k, sc0, sc1, sbs0, sbs1,
BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr,
NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,
OK: tl.constexpr,
num_warps: tl.constexpr, num_stages: tl.constexpr,
waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
cache_modifier: tl.constexpr,
):
tl.assume(sq0 > 0); tl.assume(sq1 > 0)
tl.assume(ss0 > 0); tl.assume(ss1 > 0)
tl.assume(sb0 > 0); tl.assume(sb1 > 0)
tl.assume(sc0 > 0); tl.assume(sc1 > 0)
tl.assume(sbs0 > 0); tl.assume(sbs1 > 0)
G: tl.constexpr = 32
HK: tl.constexpr = BLOCK_SIZE_K // 2
SK: tl.constexpr = BLOCK_SIZE_K // G
nm = tl.cdiv(M, BLOCK_SIZE_M)
nn = tl.cdiv(N, BLOCK_SIZE_N)
uid = tl.program_id(0)
ks = uid % NUM_KSPLIT
flat = uid // NUM_KSPLIT
if NUM_KSPLIT == 1:
gn = GROUP_SIZE_M * nn
gi = flat // gn
fm = gi * GROUP_SIZE_M
gs = min(nm - fm, GROUP_SIZE_M)
im = fm + (flat % gn) % gs
jn = (flat % gn) // gs
else:
im = flat // nn
jn = flat % nn
tl.assume(im >= 0); tl.assume(jn >= 0); tl.assume(ks >= 0)
if (ks * SPLITK_BLOCK_SIZE // 2) < K:
niter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, HK)
rm = (im * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
cq = ks * (SPLITK_BLOCK_SIZE // 2) + tl.arange(0, HK)
pq = q4 + rm[:, None] * sq0 + cq[None, :] * sq1
cs = ks * (SPLITK_BLOCK_SIZE // G) + tl.arange(0, SK)
ps = qsc + rm[:, None] * ss0 + cs[None, :] * ss1
sh_a = tl.arange(0, HK * 16)
sh_o = ks * (SPLITK_BLOCK_SIZE // 2) * 16 + sh_a
bn = (jn * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % (N // 16)
pb = b_p + bn[:, None] * sb0 + sh_o[None, :] * sb1
bsn = jn * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N // 32) * 32
bsk = ks * (SPLITK_BLOCK_SIZE // G) * 32 + tl.arange(0, SK * 32)
pbs = bs_p + bsn[:, None] * sbs0 + bsk[None, :] * sbs1
dot = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for ki in range(ks * niter, (ks + 1) * niter):
if OK:
va = tl.load(pq, cache_modifier=cache_modifier)
vas = tl.load(ps, cache_modifier=cache_modifier)
else:
lo = (ki - ks * niter) * HK
rem = K - (ks * (SPLITK_BLOCK_SIZE // 2) + lo)
va = tl.load(pq, mask=tl.arange(0, HK)[None, :] < rem, other=0, cache_modifier=cache_modifier)
sr = (2 * K) // G - (ks * (SPLITK_BLOCK_SIZE // G) + (ki - ks * niter) * SK)
vas = tl.load(ps, mask=tl.arange(0, SK)[None, :] < sr, other=0, cache_modifier=cache_modifier)
ws = (tl.load(pbs, cache_modifier=cache_modifier)
.reshape(BLOCK_SIZE_N // 32, SK // 8, 4, 16, 2, 2, 1)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_SIZE_N, SK))
if OK:
vb = tl.load(pb, cache_modifier=cache_modifier)
else:
vb = tl.load(pb,
mask=sh_a[None, :] < ((K - (ks * (SPLITK_BLOCK_SIZE // 2) + (ki - ks * niter) * HK)) * 16),
other=0, cache_modifier=cache_modifier)
vb = (vb.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_SIZE_N, HK)
.trans(1, 0))
dot = tl.dot_scaled(va, vas, "e2m1", vb, ws, "e2m1", dot)
pq += HK * sq1
ps += SK * ss1
pb += HK * 16 * sb1
pbs += BLOCK_SIZE_K * sbs1
out = dot.to(c_p.type.element_ty)
om = im * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
on = jn * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
tl.store(c_p + sc0 * om[:, None] + sc1 * on[None, :] + ks * sc_k,
out, mask=(om[:, None] < M) & (on[None, :] < N))
def _weight_views(ws, wsc, n, kp):
return ws.view(torch.uint8).reshape(n // 16, kp * 16), wsc.view(torch.uint8)
def _bufs(m, n, ns, dev):
t = (m, n, ns)
if t not in _out_pool:
y = torch.empty((m, n), dtype=torch.bfloat16, device=dev)
pp = torch.empty((ns, m, n), dtype=torch.float32, device=dev) if ns > 1 else None
_out_pool[t] = (y, pp)
return _out_pool[t]
def _setup(m, n, k, dev):
t = (m, n, k)
if t not in _grid_memo:
kp = k // 2
c = _expand(PER_SHAPE.get(t, FALLBACK).copy(), kp)
ns = c["NUM_KSPLIT"]
y, pp = _bufs(m, n, ns, dev)
g = (ns * triton.cdiv(m, c["BLOCK_SIZE_M"]) * triton.cdiv(n, c["BLOCK_SIZE_N"]),)
ck, cm, cn = (0, y.stride(0), y.stride(1)) if ns == 1 else (pp.stride(0), pp.stride(1), pp.stride(2))
info = dict(c=c, kp=kp, g=g, ns=ns, ck=ck, cm=cm, cn=cn)
if ns > 1:
info["rg"] = (triton.cdiv(m, 16), triton.cdiv(n, 64))
info["rk"] = triton.cdiv(kp, c["SPLITK_BLOCK_SIZE"] // 2)
info["pk"] = triton.next_power_of_2(ns)
_grid_memo[t] = info
return _grid_memo[t]
def _exec_fused(a, ws, wsc, m, n, k):
p = _setup(m, n, k, a.device)
y, pp = _bufs(m, n, p["ns"], a.device)
bw, bsc = _weight_views(ws, wsc, n, p["kp"])
_matmul_fused[p["g"]](
a, bw, y if p["ns"] == 1 else pp, bsc,
m, n, p["kp"],
a.stride(0), a.stride(1), bw.stride(0), bw.stride(1),
p["ck"], p["cm"], p["cn"], bsc.stride(0), bsc.stride(1),
**p["c"],
)
if p["ns"] > 1:
_merge[p["rg"]](pp, y, m, n,
pp.stride(0), pp.stride(1), pp.stride(2), y.stride(0), y.stride(1),
16, 64, p["rk"], p["pk"])
return y
def _exec_sep(a, ws, wsc, m, n, k):
kp = k // 2
QR, QC = 16, 256
fp4 = torch.empty((m, kp), dtype=torch.uint8, device=a.device)
sc = torch.empty((m, k // 32), dtype=torch.uint8, device=a.device)
_encode_block[(triton.cdiv(m, QR), triton.cdiv(k, QC))](
a, fp4, sc, m, k,
a.stride(0), a.stride(1), fp4.stride(0), fp4.stride(1), sc.stride(0), sc.stride(1),
QR, QC)
c = _expand(SEP_SHAPE.get((m, n, k), SEP_FALLBACK).copy(), kp)
ns = c["NUM_KSPLIT"]
y = torch.empty((m, n), dtype=torch.bfloat16, device=a.device)
pp = torch.empty((ns, m, n), dtype=torch.float32, device=a.device) if ns > 1 else None
bw, bsc = _weight_views(ws, wsc, n, kp)
grid_fn = lambda META: (META["NUM_KSPLIT"] * triton.cdiv(m, META["BLOCK_SIZE_M"]) * triton.cdiv(n, META["BLOCK_SIZE_N"]),)
_matmul_sep[grid_fn](
fp4, sc, bw, y if ns == 1 else pp, bsc,
m, n, kp,
fp4.stride(0), fp4.stride(1), sc.stride(0), sc.stride(1),
bw.stride(0), bw.stride(1),
0 if ns == 1 else pp.stride(0),
y.stride(0) if ns == 1 else pp.stride(1),
y.stride(1) if ns == 1 else pp.stride(2),
bsc.stride(0), bsc.stride(1), **c)
if ns > 1:
real = triton.cdiv(kp, c["SPLITK_BLOCK_SIZE"] // 2)
_merge[(triton.cdiv(m, 16), triton.cdiv(n, 64))](
pp, y, m, n,
pp.stride(0), pp.stride(1), pp.stride(2), y.stride(0), y.stride(1),
16, 64, real, triton.next_power_of_2(ns))
return y
def custom_kernel(data: input_t) -> output_t:
inp = data[0]
return _exec_fused(inp, data[3], data[4], inp.shape[0], data[1].shape[0], inp.shape[1])
scrolls · 414 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