submission 754549
flower2123 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 323 lines, June 9 Researcher Reciprocity License v1.0.
submission_v3_flower_mm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-754549?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:996bca8c64910916ff537966c9145d2ef91e7405727e1da629b5b1625b8bb5e4
license declaredunknown
license concludedunknown
authorsflower2123
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
split-k
and (a["SPLITK_BLOCK_SIZE"] % a["BLOCK_SIZE_K"] == 0)Kernel source
submission_v3_flower_mm.py323 lines
# Author: flower
import torch
import triton
import triton.language as tl
from task import input_t, output_t
_y_store = {}
_c_store = {}
_l_store = {}
TUNE = {
(4, 2880, 512): {"BLOCK_SIZE_M": 4, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 1, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
(16, 2112, 7168): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 14},
(32, 4096, 512): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
(32, 2880, 512): {"BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "cache_modifier": None, "NUM_KSPLIT": 1},
(64, 7168, 2048): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
(256, 3072, 1536): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 256, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "num_warps": 8, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "cache_modifier": None, "NUM_KSPLIT": 1},
}
TUNE_DEF = {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "num_warps": 2, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1}
SEP_TUNE = {
(32, 4096, 512): {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
(32, 2880, 512): {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},
(64, 7168, 2048): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 2},
}
SEP_TUNE_DEF = {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1}
def _fix_ksplit(half, bk, nk):
sb = triton.cdiv(2 * triton.cdiv(half, nk), bk) * bk
while nk > 1 and bk > 16:
if half % (sb // 2) == 0 and sb % bk == 0 and half % (bk // 2) == 0:
break
if half % (sb // 2) != 0 and nk > 1:
nk //= 2
elif sb % bk != 0:
nk = nk // 2 if nk > 1 else nk
if nk <= 1 and bk > 16: bk //= 2
elif half % (bk // 2) != 0 and bk > 16:
bk //= 2
else:
break
sb = triton.cdiv(2 * triton.cdiv(half, nk), bk) * bk
return sb, bk, triton.cdiv(half, sb // 2)
@triton.jit
def _quantize_mxfp4(
v, DM: tl.constexpr, DK: tl.constexpr,
):
QG: tl.constexpr = 32
NQ: tl.constexpr = DK // QG
w = v.to(tl.float32).reshape(DM, NQ, QG)
pk = tl.max(tl.abs(w), axis=-1, keep_dims=True)
pk = (pk.to(tl.int32, bitcast=True) + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
lg = ((pk >> 23) & 0xFF).to(tl.int32) - 127
ub = tl.minimum(tl.maximum(lg - 2, -127), 127)
se = ub.to(tl.uint8) + 127
hb = (ub.to(tl.int32) + 127).to(tl.uint32) << 23
hf = hb.to(tl.float32, bitcast=True)
hx = tl.broadcast_to(hf, (DM, NQ, QG)).reshape(DM, DK)
pv = hx.reshape(DM, DK // 2, 2)
ev, _ = tl.split(pv)
ev = ev.reshape(DM, DK // 2)
u16 = v.to(tl.uint16, bitcast=True).reshape(DM, DK // 2, 2)
lo, hi = tl.split(u16)
u32 = lo.to(tl.uint32) | (hi.to(tl.uint32) << 16)
u32 = u32.reshape(DM, DK // 2)
r = tl.inline_asm_elementwise(
"v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2", "=v, v, v",
[u32, ev], dtype=tl.uint32, is_pure=True, pack=1)
return (r & 0xFF).to(tl.uint8).reshape(DM, DK // 2), se.reshape(DM, NQ)
@triton.jit
def _quant_block_kernel(
src_p, fp4_p, sc_p, nrow, ncol,
s0, s1, q0, q1, c0, c1,
BM: tl.constexpr, BK: tl.constexpr,
):
ri = tl.program_id(0) * BM + tl.arange(0, BM)
ci = tl.program_id(1) * BK + tl.arange(0, BK)
d = tl.load(src_p + ri[:, None] * s0 + ci[None, :] * s1,
mask=(ri[:, None] < nrow) & (ci[None, :] < ncol), other=0.0)
f4, sc = _quantize_mxfp4(d, BM, BK)
HK: tl.constexpr = BK // 2
hc = tl.program_id(1) * HK + tl.arange(0, HK)
tl.store(fp4_p + ri[:, None] * q0 + hc[None, :] * q1,
f4, mask=(ri[:, None] < nrow) & (hc[None, :] < ncol // 2))
SK: tl.constexpr = BK // 32
si = tl.program_id(1) * SK + tl.arange(0, SK)
tl.store(sc_p + ri[:, None] * c0 + si[None, :] * c1,
sc, mask=(ri[:, None] < nrow) & (si[None, :] < ncol // 32))
@triton.heuristics({"EVEN_K": 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 _gemm_inline_quant(
ap, bp, cp, bsp, M, N, K,
sa0, sa1, sb0, sb1, sck, 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,
EVEN_K: 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)
SG: 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; pid = uid // NUM_KSPLIT
if NUM_KSPLIT == 1:
gn = GROUP_SIZE_M * nn; gi = pid // gn; fm = gi * GROUP_SIZE_M
gs = min(nm - fm, GROUP_SIZE_M); im = fm + (pid % gn) % gs; jn = (pid % gn) // gs
else:
im = pid // nn; jn = pid % nn
tl.assume(im >= 0); tl.assume(jn >= 0); tl.assume(sk >= 0)
if (sk * SPLITK_BLOCK_SIZE // 2) < K:
ni = 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 = ap + rm[:, None] * sa0 + ck[None, :] * sa1
sha = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
sho = sk * (SPLITK_BLOCK_SIZE // 2) * 16 + sha
bn = (jn * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % (N // 16)
pb = bp + bn[:, None] * sb0 + sho[None, :] * sb1
bsn = jn * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N // 32) * 32
bsk = sk * (SPLITK_BLOCK_SIZE // SG) * 32 + tl.arange(0, BLOCK_SIZE_K // SG * 32)
pbs = bsp + bsn[:, None] * sbs0 + bsk[None, :] * sbs1
acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for ki in range(sk * ni, (sk + 1) * ni):
if EVEN_K:
va = tl.load(pa); vbs = tl.load(pbs, cache_modifier=cache_modifier)
vb = tl.load(pb, cache_modifier=cache_modifier)
else:
lo = (ki - sk * ni) * 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=sha[None, :] < ((K - (sk * (SPLITK_BLOCK_SIZE // 2) + (ki - sk * ni) * (BLOCK_SIZE_K // 2))) * 16), other=0, cache_modifier=cache_modifier)
aq, asc = _quantize_mxfp4(va, BLOCK_SIZE_M, BLOCK_SIZE_K)
ws = vbs.reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // SG // 8, 4, 16, 2, 2, 1).permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SG)
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)
acc = tl.dot_scaled(aq, asc, "e2m1", bd, ws, "e2m1", acc)
pa += BLOCK_SIZE_K * sa1; pb += (BLOCK_SIZE_K // 2) * 16 * sb1; pbs += BLOCK_SIZE_K * sbs1
res = acc.to(cp.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(cp + sc0 * om[:, None] + sc1 * on[None, :] + sk * sck, res, mask=(om[:, None] < M) & (on[None, :] < N))
@triton.jit
def _accum(inp, out, M, N, isk, ism, isn, osm, osn,
TM: tl.constexpr, TN: tl.constexpr, RK: tl.constexpr, PK: tl.constexpr):
rm = (tl.program_id(0) * TM + tl.arange(0, TM)) % M
rn = (tl.program_id(1) * TN + tl.arange(0, TN)) % N
b = inp + rm[:, None] * ism + rn[None, :] * isn
s = tl.load(b).to(tl.float32)
for j in tl.static_range(1, PK):
if j < RK:
s += tl.load(b + j * isk).to(tl.float32)
tl.store(out + rm[:, None] * osm + rn[None, :] * osn, s.to(out.type.element_ty))
@triton.heuristics({"EVEN_K": 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 _gemm_preq(
q4p, scp, bp, cp, bsp, M, N, K,
sq0, sq1, ss0, ss1, sb0, sb1, sck, 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,
EVEN_K: 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)
SG: tl.constexpr = 32; HK: tl.constexpr = BLOCK_SIZE_K // 2; SCK: tl.constexpr = BLOCK_SIZE_K // SG
nm = tl.cdiv(M, BLOCK_SIZE_M); nn = tl.cdiv(N, BLOCK_SIZE_N)
uid = tl.program_id(0); sk = uid % NUM_KSPLIT; pid = uid // NUM_KSPLIT
if NUM_KSPLIT == 1:
gn = GROUP_SIZE_M * nn; gi = pid // gn; fm = gi * GROUP_SIZE_M
gs = min(nm - fm, GROUP_SIZE_M); im = fm + (pid % gn) % gs; jn = (pid % gn) // gs
else:
im = pid // nn; jn = pid % nn
tl.assume(im >= 0); tl.assume(jn >= 0); tl.assume(sk >= 0)
if (sk * SPLITK_BLOCK_SIZE // 2) < K:
ni = tl.cdiv(SPLITK_BLOCK_SIZE // 2, HK)
rm = (im * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
cq = sk * (SPLITK_BLOCK_SIZE // 2) + tl.arange(0, HK)
pq = q4p + rm[:, None] * sq0 + cq[None, :] * sq1
cs = sk * (SPLITK_BLOCK_SIZE // SG) + tl.arange(0, SCK)
ps = scp + rm[:, None] * ss0 + cs[None, :] * ss1
sha = tl.arange(0, HK * 16); sho = sk * (SPLITK_BLOCK_SIZE // 2) * 16 + sha
bn = (jn * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % (N // 16)
pb = bp + bn[:, None] * sb0 + sho[None, :] * sb1
bsn = jn * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N // 32) * 32
bsk = sk * (SPLITK_BLOCK_SIZE // SG) * 32 + tl.arange(0, SCK * 32)
pbs = bsp + bsn[:, None] * sbs0 + bsk[None, :] * sbs1
acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for ki in range(sk * ni, (sk + 1) * ni):
if EVEN_K:
va = tl.load(pq, cache_modifier=cache_modifier); vas = tl.load(ps, cache_modifier=cache_modifier)
else:
lo = (ki - sk * ni) * HK; rem = K - (sk * (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) // SG - (sk * (SPLITK_BLOCK_SIZE // SG) + (ki - sk * ni) * SCK)
vas = tl.load(ps, mask=tl.arange(0, SCK)[None, :] < sr, other=0, cache_modifier=cache_modifier)
ws = tl.load(pbs, cache_modifier=cache_modifier).reshape(BLOCK_SIZE_N // 32, SCK // 8, 4, 16, 2, 2, 1).permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_SIZE_N, SCK)
if EVEN_K:
vb = tl.load(pb, cache_modifier=cache_modifier)
else:
vb = tl.load(pb, mask=sha[None, :] < ((K - (sk * (SPLITK_BLOCK_SIZE // 2) + (ki - sk * ni) * 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)
acc = tl.dot_scaled(va, vas, "e2m1", vb, ws, "e2m1", acc)
pq += HK * sq1; ps += SCK * ss1; pb += HK * 16 * sb1; pbs += BLOCK_SIZE_K * sbs1
res = acc.to(cp.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(cp + sc0 * om[:, None] + sc1 * on[None, :] + sk * sck, res, mask=(om[:, None] < M) & (on[None, :] < N))
def _w_view(ws, wsc, n, kp):
return ws.view(torch.uint8).reshape(n // 16, kp * 16), wsc.view(torch.uint8)
def _ybuf(m, n, ns, d):
k = (m, n, ns)
if k not in _y_store:
_y_store[k] = (torch.empty((m, n), dtype=torch.bfloat16, device=d),
torch.empty((ns, m, n), dtype=torch.float32, device=d) if ns > 1 else None)
return _y_store[k]
def _resolve(m, n, k):
t = (m, n, k)
if t not in _c_store:
c = TUNE.get(t, TUNE_DEF).copy()
kp = k // 2
if c["NUM_KSPLIT"] > 1:
sb, bk, ns = _fix_ksplit(kp, c["BLOCK_SIZE_K"], c["NUM_KSPLIT"])
c["SPLITK_BLOCK_SIZE"], c["BLOCK_SIZE_K"], c["NUM_KSPLIT"] = sb, bk, ns
else:
c["SPLITK_BLOCK_SIZE"], c["NUM_KSPLIT"] = 2 * kp, 1
if c["BLOCK_SIZE_K"] >= 2 * kp:
c["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * kp)
c["SPLITK_BLOCK_SIZE"], c["NUM_KSPLIT"] = 2 * kp, 1
c["BLOCK_SIZE_N"] = max(c["BLOCK_SIZE_N"], 32)
_c_store[t] = c
return _c_store[t]
def _pre(m, n, k, d):
t = (m, n, k)
if t not in _l_store:
c = _resolve(m, n, k); kp = k // 2; ns = c["NUM_KSPLIT"]
y, pp = _ybuf(m, n, ns, d)
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))
r = dict(c=c, kp=kp, g=g, ns=ns, ck=ck, cm=cm, cn=cn)
if ns > 1:
r["rg"] = (triton.cdiv(m, 16), triton.cdiv(n, 64))
r["rk"] = triton.cdiv(kp, c["SPLITK_BLOCK_SIZE"] // 2)
r["pk"] = triton.next_power_of_2(ns)
_l_store[t] = r
return _l_store[t]
def _go_fused(a, ws, wsc, m, n, k):
p = _pre(m, n, k, a.device)
y, pp = _ybuf(m, n, p["ns"], a.device)
bw, bs = _w_view(ws, wsc, n, p["kp"])
_gemm_inline_quant[p["g"]](
a, bw, y if p["ns"] == 1 else pp, bs, m, n, p["kp"],
a.stride(0), a.stride(1), bw.stride(0), bw.stride(1),
p["ck"], p["cm"], p["cn"], bs.stride(0), bs.stride(1), **p["c"])
if p["ns"] > 1:
_accum[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 _go_sep(a, ws, wsc, m, n, k):
kp, QM, QK = k // 2, 16, 256
f4 = torch.empty((m, kp), dtype=torch.uint8, device=a.device)
sc = torch.empty((m, k // 32), dtype=torch.uint8, device=a.device)
_quant_block_kernel[(triton.cdiv(m, QM), triton.cdiv(k, QK))](
a, f4, sc, m, k, a.stride(0), a.stride(1), f4.stride(0), f4.stride(1), sc.stride(0), sc.stride(1), QM, QK)
c = SEP_TUNE.get((m, n, k), SEP_TUNE_DEF).copy()
if c["NUM_KSPLIT"] > 1:
sb, bk, ns = _fix_ksplit(kp, c["BLOCK_SIZE_K"], c["NUM_KSPLIT"])
c["SPLITK_BLOCK_SIZE"], c["BLOCK_SIZE_K"], c["NUM_KSPLIT"] = sb, bk, ns
else:
c["SPLITK_BLOCK_SIZE"], c["NUM_KSPLIT"] = 2 * kp, 1
if c["BLOCK_SIZE_K"] >= 2 * kp:
c["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * kp); c["SPLITK_BLOCK_SIZE"], c["NUM_KSPLIT"] = 2 * kp, 1
c["BLOCK_SIZE_N"] = max(c["BLOCK_SIZE_N"], 32)
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, bs = _w_view(ws, wsc, n, kp)
gf = lambda META: (META["NUM_KSPLIT"] * triton.cdiv(m, META["BLOCK_SIZE_M"]) * triton.cdiv(n, META["BLOCK_SIZE_N"]),)
_gemm_preq[gf](f4, sc, bw, y if ns == 1 else pp, bs, m, n, kp,
f4.stride(0), f4.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), bs.stride(0), bs.stride(1), **c)
if ns > 1:
rk = triton.cdiv(kp, c["SPLITK_BLOCK_SIZE"] // 2)
_accum[(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, rk, triton.next_power_of_2(ns))
return y
def custom_kernel(data: input_t) -> output_t:
x = data[0]
return _go_fused(x, data[3], data[4], x.shape[0], data[1].shape[0], x.shape[1])
scrolls · 323 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 754367.
+ # Author: flower+import torchimport tritonimport triton.language as tlfrom task import input_t, output_t- _out_pool = {}- _params = {}- _grid_memo = {}+ _y_store = {}+ _c_store = {}+ _l_store = {}-- 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),+ TUNE = {+ (4, 2880, 512): {"BLOCK_SIZE_M": 4, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 1, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},+ (16, 2112, 7168): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 14},+ (32, 4096, 512): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},+ (32, 2880, 512): {"BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "cache_modifier": None, "NUM_KSPLIT": 1},+ (64, 7168, 2048): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},+ (256, 3072, 1536): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 256, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 1, "num_warps": 8, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "cache_modifier": None, "NUM_KSPLIT": 1},}- FALLBACK = dict(tm=16, tn=32, tk=256, gm=1, warp=2, pipe=2, occ=0, kdim=16, cmod=".cg", ks=1)+ TUNE_DEF = {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 1, "num_warps": 2, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 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_TUNE = {+ (32, 4096, 512): {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},+ (32, 2880, 512): {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 1},+ (64, 7168, 2048): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 2},}- SEP_FALLBACK = dict(tm=16, tn=64, tk=256, gm=4, warp=4, pipe=2, occ=0, kdim=16, cmod=".cg", ks=1)+ SEP_TUNE_DEF = {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256, "GROUP_SIZE_M": 4, "num_warps": 4, "num_stages": 2, "waves_per_eu": 0, "matrix_instr_nonkdim": 16, "cache_modifier": ".cg", "NUM_KSPLIT": 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:+ def _fix_ksplit(half, bk, nk):+ sb = triton.cdiv(2 * triton.cdiv(half, nk), bk) * bk+ while nk > 1 and bk > 16:+ if half % (sb // 2) == 0 and sb % bk == 0 and half % (bk // 2) == 0:break- if kp % (sb // 2) != 0 and ns > 1:- ns //= 2+ if half % (sb // 2) != 0 and nk > 1:+ nk //= 2elif 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:+ nk = nk // 2 if nk > 1 else nk+ if nk <= 1 and bk > 16: bk //= 2+ elif half % (bk // 2) != 0 and bk > 16:bk //= 2else:break- sb = triton.cdiv(2 * triton.cdiv(kp, ns), bk) * bk- return sb, bk, triton.cdiv(kp, sb // 2)+ sb = triton.cdiv(2 * triton.cdiv(half, nk), bk) * bk+ return sb, bk, triton.cdiv(half, 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,+ def _quantize_mxfp4(+ v, DM: tl.constexpr, DK: 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+ QG: tl.constexpr = 32+ NQ: tl.constexpr = DK // QG+ w = v.to(tl.float32).reshape(DM, NQ, QG)+ pk = tl.max(tl.abs(w), axis=-1, keep_dims=True)+ pk = (pk.to(tl.int32, bitcast=True) + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000+ lg = ((pk >> 23) & 0xFF).to(tl.int32) - 127+ ub = tl.minimum(tl.maximum(lg - 2, -127), 127)+ se = ub.to(tl.uint8) + 127hb = (ub.to(tl.int32) + 127).to(tl.uint32) << 23hf = 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)+ hx = tl.broadcast_to(hf, (DM, NQ, QG)).reshape(DM, DK)+ pv = hx.reshape(DM, DK // 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(+ ev = ev.reshape(DM, DK // 2)+ u16 = v.to(tl.uint16, bitcast=True).reshape(DM, DK // 2, 2)+ lo, hi = tl.split(u16)+ u32 = lo.to(tl.uint32) | (hi.to(tl.uint32) << 16)+ u32 = u32.reshape(DM, DK // 2)+ r = 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)+ [u32, ev], dtype=tl.uint32, is_pure=True, pack=1)+ return (r & 0xFF).to(tl.uint8).reshape(DM, DK // 2), se.reshape(DM, NQ)@triton.jit- def _encode_block(- inp, o_fp4, o_sc,- nrow, ncol,- si0, si1, sq0, sq1, ss0, ss1,- TR: tl.constexpr, TC: tl.constexpr,+ def _quant_block_kernel(+ src_p, fp4_p, sc_p, nrow, ncol,+ s0, s1, q0, q1, c0, c1,+ BM: tl.constexpr, BK: 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))+ ri = tl.program_id(0) * BM + tl.arange(0, BM)+ ci = tl.program_id(1) * BK + tl.arange(0, BK)+ d = tl.load(src_p + ri[:, None] * s0 + ci[None, :] * s1,+ mask=(ri[:, None] < nrow) & (ci[None, :] < ncol), other=0.0)+ f4, sc = _quantize_mxfp4(d, BM, BK)+ HK: tl.constexpr = BK // 2+ hc = tl.program_id(1) * HK + tl.arange(0, HK)+ tl.store(fp4_p + ri[:, None] * q0 + hc[None, :] * q1,+ f4, mask=(ri[:, None] < nrow) & (hc[None, :] < ncol // 2))+ SK: tl.constexpr = BK // 32+ si = tl.program_id(1) * SK + tl.arange(0, SK)+ tl.store(sc_p + ri[:, None] * c0 + si[None, :] * c1,+ sc, mask=(ri[:, None] < nrow) & (si[None, :] < ncol // 32))- @triton.heuristics({"OK": lambda a: (a["K"] % (a["BLOCK_SIZE_K"] // 2) == 0)+ @triton.heuristics({"EVEN_K": 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,+ def _gemm_inline_quant(+ ap, bp, cp, bsp, M, N, K,+ sa0, sa1, sb0, sb1, sck, 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,+ GROUP_SIZE_M: tl.constexpr, NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,+ EVEN_K: tl.constexpr,num_warps: tl.constexpr, num_stages: tl.constexpr,- waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,- cache_modifier: 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+ 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)+ SG: 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; pid = uid // NUM_KSPLITif 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+ gn = GROUP_SIZE_M * nn; gi = pid // gn; fm = gi * GROUP_SIZE_M+ gs = min(nm - fm, GROUP_SIZE_M); im = fm + (pid % gn) % gs; jn = (pid % gn) // gselse:- im = flat // nn- jn = flat % nn+ im = pid // nn; jn = pid % nntl.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)+ ni = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)rm = (im * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % Mck = 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+ pa = ap + rm[:, None] * sa0 + ck[None, :] * sa1+ sha = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)+ sho = sk * (SPLITK_BLOCK_SIZE // 2) * 16 + shabn = (jn * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % (N // 16)- pb = b_p + bn[:, None] * sb0 + sh_o[None, :] * sb1-+ pb = bp + bn[:, None] * sb0 + sho[None, :] * sb1bsn = 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)+ bsk = sk * (SPLITK_BLOCK_SIZE // SG) * 32 + tl.arange(0, BLOCK_SIZE_K // SG * 32)+ pbs = bsp + bsn[:, None] * sbs0 + bsk[None, :] * sbs1+ acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)+ for ki in range(sk * ni, (sk + 1) * ni):+ if EVEN_K:+ 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+ lo = (ki - sk * ni) * BLOCK_SIZE_Kva = 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)+ vb = tl.load(pb, mask=sha[None, :] < ((K - (sk * (SPLITK_BLOCK_SIZE // 2) + (ki - sk * ni) * (BLOCK_SIZE_K // 2))) * 16), other=0, cache_modifier=cache_modifier)+ aq, asc = _quantize_mxfp4(va, BLOCK_SIZE_M, BLOCK_SIZE_K)+ ws = vbs.reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // SG // 8, 4, 16, 2, 2, 1).permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SG)+ 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)+ acc = tl.dot_scaled(aq, asc, "e2m1", bd, ws, "e2m1", acc)+ pa += BLOCK_SIZE_K * sa1; pb += (BLOCK_SIZE_K // 2) * 16 * sb1; pbs += BLOCK_SIZE_K * sbs1+ res = acc.to(cp.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))+ tl.store(cp + sc0 * om[:, None] + sc1 * on[None, :] + sk * sck, res, 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):+ def _accum(inp, out, M, N, isk, ism, isn, osm, osn,+ TM: tl.constexpr, TN: tl.constexpr, RK: tl.constexpr, PK: tl.constexpr):rm = (tl.program_id(0) * TM + tl.arange(0, TM)) % Mrn = (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))+ b = inp + rm[:, None] * ism + rn[None, :] * isn+ s = tl.load(b).to(tl.float32)+ for j in tl.static_range(1, PK):+ if j < RK:+ s += tl.load(b + j * isk).to(tl.float32)+ tl.store(out + rm[:, None] * osm + rn[None, :] * osn, s.to(out.type.element_ty))- @triton.heuristics({"OK": lambda a: (a["K"] % (a["BLOCK_SIZE_K"] // 2) == 0)+ @triton.heuristics({"EVEN_K": 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,+ def _gemm_preq(+ q4p, scp, bp, cp, bsp, M, N, K,+ sq0, sq1, ss0, ss1, sb0, sb1, sck, 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,+ GROUP_SIZE_M: tl.constexpr, NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,+ EVEN_K: tl.constexpr,num_warps: tl.constexpr, num_stages: tl.constexpr,- waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,- cache_modifier: 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(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+ SG: tl.constexpr = 32; HK: tl.constexpr = BLOCK_SIZE_K // 2; SCK: tl.constexpr = BLOCK_SIZE_K // SG+ nm = tl.cdiv(M, BLOCK_SIZE_M); nn = tl.cdiv(N, BLOCK_SIZE_N)+ uid = tl.program_id(0); sk = uid % NUM_KSPLIT; pid = uid // NUM_KSPLITif 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+ gn = GROUP_SIZE_M * nn; gi = pid // gn; fm = gi * GROUP_SIZE_M+ gs = min(nm - fm, GROUP_SIZE_M); im = fm + (pid % gn) % gs; jn = (pid % gn) // gselse:- 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)+ im = pid // nn; jn = pid % nn+ tl.assume(im >= 0); tl.assume(jn >= 0); tl.assume(sk >= 0)+ if (sk * SPLITK_BLOCK_SIZE // 2) < K:+ ni = 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+ cq = sk * (SPLITK_BLOCK_SIZE // 2) + tl.arange(0, HK)+ pq = q4p + rm[:, None] * sq0 + cq[None, :] * sq1+ cs = sk * (SPLITK_BLOCK_SIZE // SG) + tl.arange(0, SCK)+ ps = scp + rm[:, None] * ss0 + cs[None, :] * ss1+ sha = tl.arange(0, HK * 16); sho = sk * (SPLITK_BLOCK_SIZE // 2) * 16 + shabn = (jn * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % (N // 16)- pb = b_p + bn[:, None] * sb0 + sh_o[None, :] * sb1-+ pb = bp + bn[:, None] * sb0 + sho[None, :] * sb1bsn = 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)+ bsk = sk * (SPLITK_BLOCK_SIZE // SG) * 32 + tl.arange(0, SCK * 32)+ pbs = bsp + bsn[:, None] * sbs0 + bsk[None, :] * sbs1+ acc = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)+ for ki in range(sk * ni, (sk + 1) * ni):+ if EVEN_K:+ 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)+ lo = (ki - sk * ni) * HK; rem = K - (sk * (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:+ sr = (2 * K) // SG - (sk * (SPLITK_BLOCK_SIZE // SG) + (ki - sk * ni) * SCK)+ vas = tl.load(ps, mask=tl.arange(0, SCK)[None, :] < sr, other=0, cache_modifier=cache_modifier)+ ws = tl.load(pbs, cache_modifier=cache_modifier).reshape(BLOCK_SIZE_N // 32, SCK // 8, 4, 16, 2, 2, 1).permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_SIZE_N, SCK)+ if EVEN_K: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)+ vb = tl.load(pb, mask=sha[None, :] < ((K - (sk * (SPLITK_BLOCK_SIZE // 2) + (ki - sk * ni) * 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)+ acc = tl.dot_scaled(va, vas, "e2m1", vb, ws, "e2m1", acc)+ pq += HK * sq1; ps += SCK * ss1; pb += HK * 16 * sb1; pbs += BLOCK_SIZE_K * sbs1+ res = acc.to(cp.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))+ tl.store(cp + sc0 * om[:, None] + sc1 * on[None, :] + sk * sck, res, mask=(om[:, None] < M) & (on[None, :] < N))- def _weight_views(ws, wsc, n, kp):+ def _w_view(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 _ybuf(m, n, ns, d):+ k = (m, n, ns)+ if k not in _y_store:+ _y_store[k] = (torch.empty((m, n), dtype=torch.bfloat16, device=d),+ torch.empty((ns, m, n), dtype=torch.float32, device=d) if ns > 1 else None)+ return _y_store[k]- def _setup(m, n, k, dev):+ def _resolve(m, n, k):t = (m, n, k)- if t not in _grid_memo:+ if t not in _c_store:+ c = TUNE.get(t, TUNE_DEF).copy()kp = k // 2- c = _expand(PER_SHAPE.get(t, FALLBACK).copy(), kp)- ns = c["NUM_KSPLIT"]- y, pp = _bufs(m, n, ns, dev)+ if c["NUM_KSPLIT"] > 1:+ sb, bk, ns = _fix_ksplit(kp, c["BLOCK_SIZE_K"], c["NUM_KSPLIT"])+ c["SPLITK_BLOCK_SIZE"], c["BLOCK_SIZE_K"], c["NUM_KSPLIT"] = sb, bk, ns+ else:+ c["SPLITK_BLOCK_SIZE"], c["NUM_KSPLIT"] = 2 * kp, 1+ if c["BLOCK_SIZE_K"] >= 2 * kp:+ c["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * kp)+ c["SPLITK_BLOCK_SIZE"], c["NUM_KSPLIT"] = 2 * kp, 1+ c["BLOCK_SIZE_N"] = max(c["BLOCK_SIZE_N"], 32)+ _c_store[t] = c+ return _c_store[t]+++ def _pre(m, n, k, d):+ t = (m, n, k)+ if t not in _l_store:+ c = _resolve(m, n, k); kp = k // 2; ns = c["NUM_KSPLIT"]+ y, pp = _ybuf(m, n, ns, d)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)+ r = 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]+ r["rg"] = (triton.cdiv(m, 16), triton.cdiv(n, 64))+ r["rk"] = triton.cdiv(kp, c["SPLITK_BLOCK_SIZE"] // 2)+ r["pk"] = triton.next_power_of_2(ns)+ _l_store[t] = r+ return _l_store[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"],+ def _go_fused(a, ws, wsc, m, n, k):+ p = _pre(m, n, k, a.device)+ y, pp = _ybuf(m, n, p["ns"], a.device)+ bw, bs = _w_view(ws, wsc, n, p["kp"])+ _gemm_inline_quant[p["g"]](+ a, bw, y if p["ns"] == 1 else pp, bs, 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"],- )+ p["ck"], p["cm"], p["cn"], bs.stride(0), bs.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"])+ _accum[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)+ def _go_sep(a, ws, wsc, m, n, k):+ kp, QM, QK = k // 2, 16, 256+ f4 = 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)+ _quant_block_kernel[(triton.cdiv(m, QM), triton.cdiv(k, QK))](+ a, f4, sc, m, k, a.stride(0), a.stride(1), f4.stride(0), f4.stride(1), sc.stride(0), sc.stride(1), QM, QK)+ c = SEP_TUNE.get((m, n, k), SEP_TUNE_DEF).copy()+ if c["NUM_KSPLIT"] > 1:+ sb, bk, ns = _fix_ksplit(kp, c["BLOCK_SIZE_K"], c["NUM_KSPLIT"])+ c["SPLITK_BLOCK_SIZE"], c["BLOCK_SIZE_K"], c["NUM_KSPLIT"] = sb, bk, ns+ else:+ c["SPLITK_BLOCK_SIZE"], c["NUM_KSPLIT"] = 2 * kp, 1+ if c["BLOCK_SIZE_K"] >= 2 * kp:+ c["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * kp); c["SPLITK_BLOCK_SIZE"], c["NUM_KSPLIT"] = 2 * kp, 1+ c["BLOCK_SIZE_N"] = max(c["BLOCK_SIZE_N"], 32)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)-+ bw, bs = _w_view(ws, wsc, n, kp)+ gf = lambda META: (META["NUM_KSPLIT"] * triton.cdiv(m, META["BLOCK_SIZE_M"]) * triton.cdiv(n, META["BLOCK_SIZE_N"]),)+ _gemm_preq[gf](f4, sc, bw, y if ns == 1 else pp, bs, m, n, kp,+ f4.stride(0), f4.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), bs.stride(0), bs.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))+ rk = triton.cdiv(kp, c["SPLITK_BLOCK_SIZE"] // 2)+ _accum[(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, rk, triton.next_power_of_2(ns))return ydef 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])+ x = data[0]+ return _go_fused(x, data[3], data[4], x.shape[0], data[1].shape[0], x.shape[1])
scrolls · 632 diff lines total
Best evidence level for this revision: reported
JSON