submission 719665
Sami · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 287 lines, June 9 Researcher Reciprocity License v1.0.
submission_v58_selective_cachemods.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-719665?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:315ed491b1b28e2af325d4b54018d91bbce96874dc232e04185b2612a82ea745
license declaredunknown
license concludedunknown
authorsSami
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MXFP4 GEMM - v58: selective cache modifiers on winning paths only.persistent-kernel
- persistent K=2048 pathsplit-k
- fused split-K path for the K=7168 shapestile-k = 512
BM=BM, BN=BN, BK=512, KI=KI, NSMS=nc, num_warps=nw, num_stages=ns, waves_per_eu=2)tile-m = 16
_reduce_kernel[(triton.cdiv(m, 16), triton.cdiv(n, BN))](p, out, m, n, p.stride(0), p.stride(1), p.stride(2), out.stride(0), out.stride(1), NS=SK, BM=16, BN=BN)tile-n = 16
else: BM, BN = 16, 128; SK = max(1, k // BK) if k >= 4096 else 1; KI = max(1, triton.cdiv(k, max(SK, 1) * BK)); nw, ns, even = 4, 1, FalseKernel source
submission_v58_selective_cachemods.py287 lines
"""
MXFP4 GEMM - v58: selective cache modifiers on winning paths only.
Only enable aiter-style cache hints where v57 helped:
- fused split-K path for the K=7168 shapes
- persistent K=2048 path
Keep the baseline kernels for K=512 fused shapes and the K=1536 lean path.
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import _mxfp4_quant_op
_u8 = torch.uint8
_bf16 = torch.bfloat16
_f32 = torch.float32
NUM_SMS = 256
@triton.jit
def _fused_gemm_kernel(
a_ptr, b_ptr, c_ptr, bs_ptr, M, N, K,
stride_am, stride_ak, stride_bk, stride_bn,
stride_ck, stride_cm, stride_cn, SN,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr, NUM_KSPLIT: tl.constexpr,
K_ITERS_PER_SPLIT: tl.constexpr, EVEN_MNK: tl.constexpr,
num_warps: tl.constexpr, num_stages: tl.constexpr, waves_per_eu: tl.constexpr,
):
K_HALF: tl.constexpr = BLOCK_K // 2; SCALE_K: tl.constexpr = BLOCK_K // 32; SN32 = SN * 32
GRID_MN = tl.cdiv(M, BLOCK_M) * tl.cdiv(N, BLOCK_N)
pu = tl.program_id(0); pk = pu // GRID_MN; pid = pu % GRID_MN
npm = tl.cdiv(M, BLOCK_M); npn = tl.cdiv(N, BLOCK_N)
if NUM_KSPLIT == 1:
g = GROUP_SIZE_M * npn; gid = pid // g; fm = gid * GROUP_SIZE_M
gsm = tl.minimum(npm - fm, GROUP_SIZE_M); pid_m = fm + ((pid % g) % gsm); pid_n = (pid % g) // gsm
else:
pid_m = pid // npn; pid_n = pid % npn
om = pid_m * BLOCK_M + tl.arange(0, BLOCK_M); on = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
sr = (on // 32) * SN32 + ((on >> 4) & 1) + (on & 15) * 4
ab = a_ptr + om[:, None] * stride_am; acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
ks = pk * K_ITERS_PER_SPLIT * BLOCK_K
for ki in range(K_ITERS_PER_SPLIT):
kst = ks + ki * BLOCK_K
if kst < K:
kh = kst // 2; ak = kst + tl.arange(0, BLOCK_K); bk = kh + tl.arange(0, K_HALF)
if EVEN_MNK:
at = tl.load(ab + ak[None, :] * stride_ak)
bt = tl.load(b_ptr + bk[:, None] * stride_bk + on[None, :] * stride_bn)
else:
at = tl.load(ab + ak[None, :] * stride_ak, mask=(om[:, None] < M) & (ak[None, :] < K), other=0.0)
bt = tl.load(b_ptr + bk[:, None] * stride_bk + on[None, :] * stride_bn, mask=(bk[:, None] < (K // 2)) & (on[None, :] < N), other=0)
sk = (kst // 32) + tl.arange(0, SCALE_K); sc = (sk // 8) * 256 + ((sk >> 2) & 1) * 2 + (sk & 3) * 64
if EVEN_MNK: bs = tl.load(bs_ptr + sr[:, None] + sc[None, :])
else: bs = tl.load(bs_ptr + sr[:, None] + sc[None, :], mask=(on[:, None] < N) & (sk[None, :] < (K // 32)), other=0)
af, asc = _mxfp4_quant_op(at, BLOCK_K, BLOCK_M, 32)
acc = tl.dot_scaled(af, asc, "e2m1", bt, bs, "e2m1", acc)
c = acc.to(tl.bfloat16) if NUM_KSPLIT == 1 else acc
cm = (om[:, None] < M) & (on[None, :] < N)
tl.store(c_ptr + pk * stride_ck + om[:, None] * stride_cm + on[None, :] * stride_cn, c, mask=cm)
@triton.jit
def _fused_cache_gemm_kernel(
a_ptr, b_ptr, c_ptr, bs_ptr, M, N, K,
stride_am, stride_ak, stride_bk, stride_bn,
stride_ck, stride_cm, stride_cn, SN,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr, NUM_KSPLIT: tl.constexpr,
K_ITERS_PER_SPLIT: tl.constexpr, EVEN_MNK: tl.constexpr,
num_warps: tl.constexpr, num_stages: tl.constexpr, waves_per_eu: tl.constexpr,
):
K_HALF: tl.constexpr = BLOCK_K // 2; SCALE_K: tl.constexpr = BLOCK_K // 32; SN32 = SN * 32
GRID_MN = tl.cdiv(M, BLOCK_M) * tl.cdiv(N, BLOCK_N)
pu = tl.program_id(0); pk = pu // GRID_MN; pid = pu % GRID_MN
npm = tl.cdiv(M, BLOCK_M); npn = tl.cdiv(N, BLOCK_N)
if NUM_KSPLIT == 1:
g = GROUP_SIZE_M * npn; gid = pid // g; fm = gid * GROUP_SIZE_M
gsm = tl.minimum(npm - fm, GROUP_SIZE_M); pid_m = fm + ((pid % g) % gsm); pid_n = (pid % g) // gsm
else:
pid_m = pid // npn; pid_n = pid % npn
om = pid_m * BLOCK_M + tl.arange(0, BLOCK_M); on = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
sr = (on // 32) * SN32 + ((on >> 4) & 1) + (on & 15) * 4
ab = a_ptr + om[:, None] * stride_am; acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
ks = pk * K_ITERS_PER_SPLIT * BLOCK_K
for ki in range(K_ITERS_PER_SPLIT):
kst = ks + ki * BLOCK_K
if kst < K:
kh = kst // 2; ak = kst + tl.arange(0, BLOCK_K); bk = kh + tl.arange(0, K_HALF)
if EVEN_MNK:
at = tl.load(ab + ak[None, :] * stride_ak)
bt = tl.load(b_ptr + bk[:, None] * stride_bk + on[None, :] * stride_bn, cache_modifier=".cg")
else:
at = tl.load(ab + ak[None, :] * stride_ak, mask=(om[:, None] < M) & (ak[None, :] < K), other=0.0)
bt = tl.load(b_ptr + bk[:, None] * stride_bk + on[None, :] * stride_bn, mask=(bk[:, None] < (K // 2)) & (on[None, :] < N), other=0, cache_modifier=".cg")
sk = (kst // 32) + tl.arange(0, SCALE_K); sc = (sk // 8) * 256 + ((sk >> 2) & 1) * 2 + (sk & 3) * 64
if EVEN_MNK:
bs = tl.load(bs_ptr + sr[:, None] + sc[None, :], cache_modifier=".cg")
else:
bs = tl.load(bs_ptr + sr[:, None] + sc[None, :], mask=(on[:, None] < N) & (sk[None, :] < (K // 32)), other=0, cache_modifier=".cg")
af, asc = _mxfp4_quant_op(at, BLOCK_K, BLOCK_M, 32)
acc = tl.dot_scaled(af, asc, "e2m1", bt, bs, "e2m1", acc)
c = acc.to(tl.bfloat16) if NUM_KSPLIT == 1 else acc
cm = (om[:, None] < M) & (on[None, :] < N)
tl.store(c_ptr + pk * stride_ck + om[:, None] * stride_cm + on[None, :] * stride_cn, c, mask=cm, cache_modifier=".wt")
@triton.jit
def _lean_gemm_kernel(
aq_ptr, bq_ptr, as_ptr, bs_ptr, c_ptr, M, N, K,
saqm, saqk, sbqk, sbqn, sasm, sask, scm, scn, BSN,
BM: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr,
GSM: tl.constexpr, KI: tl.constexpr, EVEN: tl.constexpr,
num_warps: tl.constexpr, num_stages: tl.constexpr, waves_per_eu: tl.constexpr,
):
KH: tl.constexpr = BK // 2; SK: tl.constexpr = BK // 32; BSN32 = BSN * 32
pid = tl.program_id(0); npm = tl.cdiv(M, BM); npn = tl.cdiv(N, BN)
g = GSM * npn; gid = pid // g; fm = gid * GSM
gsm = tl.minimum(npm - fm, GSM)
pm = fm + ((pid % g) % gsm); pn = (pid % g) // gsm
om = pm * BM + tl.arange(0, BM); on = pn * BN + tl.arange(0, BN)
sr = (on // 32) * BSN32 + ((on >> 4) & 1) + (on & 15) * 4
acc = tl.zeros((BM, BN), dtype=tl.float32)
for ki in range(KI):
kst = ki * BK; kh = kst // 2
ak = kh + tl.arange(0, KH); bk = kh + tl.arange(0, KH)
at = tl.load(aq_ptr + om[:, None] * saqm + ak[None, :] * saqk)
bt = tl.load(bq_ptr + bk[:, None] * sbqk + on[None, :] * sbqn)
ski = kst // 32; sk = ski + tl.arange(0, SK)
a_s = tl.load(as_ptr + om[:, None] * sasm + sk[None, :] * sask)
sc2 = (sk // 8) * 256 + ((sk >> 2) & 1) * 2 + (sk & 3) * 64
b_s = tl.load(bs_ptr + sr[:, None] + sc2[None, :])
acc = tl.dot_scaled(at, a_s, "e2m1", bt, b_s, "e2m1", acc)
c = acc.to(tl.bfloat16)
cm = (om[:, None] < M) & (on[None, :] < N)
tl.store(c_ptr + om[:, None] * scm + on[None, :] * scn, c, mask=cm)
@triton.jit
def _persistent_lean_kernel(
aq_ptr, bq_ptr, as_ptr, bs_ptr, c_ptr, M, N, K, TT,
saqm, saqk, sbqk, sbqn, sasm, sask, scm, scn, BSN,
BM: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr,
KI: tl.constexpr, NSMS: tl.constexpr,
num_warps: tl.constexpr, num_stages: tl.constexpr, waves_per_eu: tl.constexpr,
):
KH: tl.constexpr = BK // 2; SK: tl.constexpr = BK // 32; BSN32 = BSN * 32
npn = tl.cdiv(N, BN); pid = tl.program_id(0); tpc = tl.cdiv(TT, NSMS)
for _i in range(tpc):
tid = pid + _i * NSMS
if tid < TT:
pm = tid // npn; pn = tid % npn
om = pm * BM + tl.arange(0, BM); on = pn * BN + tl.arange(0, BN)
sr = (on // 32) * BSN32 + ((on >> 4) & 1) + (on & 15) * 4
acc = tl.zeros((BM, BN), dtype=tl.float32)
for ki in range(KI):
kst = ki * BK; kh = kst // 2
ak = kh + tl.arange(0, KH); bk = kh + tl.arange(0, KH)
at = tl.load(aq_ptr + om[:, None] * saqm + ak[None, :] * saqk)
bt = tl.load(bq_ptr + bk[:, None] * sbqk + on[None, :] * sbqn)
ski = kst // 32; sk = ski + tl.arange(0, SK)
a_s = tl.load(as_ptr + om[:, None] * sasm + sk[None, :] * sask)
sc2 = (sk // 8) * 256 + ((sk >> 2) & 1) * 2 + (sk & 3) * 64
b_s = tl.load(bs_ptr + sr[:, None] + sc2[None, :])
acc = tl.dot_scaled(at, a_s, "e2m1", bt, b_s, "e2m1", acc)
c = acc.to(tl.bfloat16)
cm = (om[:, None] < M) & (on[None, :] < N)
tl.store(c_ptr + om[:, None] * scm + on[None, :] * scn, c, mask=cm)
@triton.jit
def _persistent_cache_lean_kernel(
aq_ptr, bq_ptr, as_ptr, bs_ptr, c_ptr, M, N, K, TT,
saqm, saqk, sbqk, sbqn, sasm, sask, scm, scn, BSN,
BM: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr,
KI: tl.constexpr, NSMS: tl.constexpr,
num_warps: tl.constexpr, num_stages: tl.constexpr, waves_per_eu: tl.constexpr,
):
KH: tl.constexpr = BK // 2; SK: tl.constexpr = BK // 32; BSN32 = BSN * 32
npn = tl.cdiv(N, BN); pid = tl.program_id(0); tpc = tl.cdiv(TT, NSMS)
for _i in range(tpc):
tid = pid + _i * NSMS
if tid < TT:
pm = tid // npn; pn = tid % npn
om = pm * BM + tl.arange(0, BM); on = pn * BN + tl.arange(0, BN)
sr = (on // 32) * BSN32 + ((on >> 4) & 1) + (on & 15) * 4
acc = tl.zeros((BM, BN), dtype=tl.float32)
for ki in range(KI):
kst = ki * BK; kh = kst // 2
ak = kh + tl.arange(0, KH); bk = kh + tl.arange(0, KH)
at = tl.load(aq_ptr + om[:, None] * saqm + ak[None, :] * saqk)
bt = tl.load(bq_ptr + bk[:, None] * sbqk + on[None, :] * sbqn, cache_modifier=".cg")
ski = kst // 32; sk = ski + tl.arange(0, SK)
a_s = tl.load(as_ptr + om[:, None] * sasm + sk[None, :] * sask)
sc2 = (sk // 8) * 256 + ((sk >> 2) & 1) * 2 + (sk & 3) * 64
b_s = tl.load(bs_ptr + sr[:, None] + sc2[None, :], cache_modifier=".cg")
acc = tl.dot_scaled(at, a_s, "e2m1", bt, b_s, "e2m1", acc)
c = acc.to(tl.bfloat16)
cm = (om[:, None] < M) & (on[None, :] < N)
tl.store(c_ptr + om[:, None] * scm + on[None, :] * scn, c, mask=cm, cache_modifier=".wt")
@triton.jit
def _reduce_kernel(pp, op, M, N, spk, spm, spn, som, son,
NS: tl.constexpr, BM: tl.constexpr, BN: tl.constexpr):
pm = tl.program_id(0) * BM + tl.arange(0, BM); pn = tl.program_id(1) * BN + tl.arange(0, BN)
m = (pm[:, None] < M) & (pn[None, :] < N); a = tl.zeros((BM, BN), dtype=tl.float32)
for k in range(NS):
a += tl.load(pp + k * spk + pm[:, None] * spm + pn[None, :] * spn, mask=m, other=0.0)
tl.store(op + pm[:, None] * som + pn[None, :] * son, a.to(tl.bfloat16), mask=m)
_FUSED_CFGS = {
(4, 2880, 512): (16, 64, 1, 1, 4, 1, False),
(32, 4096, 512): (16, 64, 1, 1, 4, 1, True),
(32, 2880, 512): (16, 64, 1, 1, 4, 1, True),
(256, 2880, 512): (16, 64, 1, 1, 4, 1, True),
(8, 2112, 7168): (16, 64, 14, 1, 4, 1, False),
(16, 2112, 7168): (16, 64, 14, 1, 4, 1, True),
}
_PERSISTENT_CFGS = {(64, 7168, 2048): (16, 128, 4, 4, 2)}
_LEAN_CFGS = {
(256, 3072, 1536): (16, 128, 3, 4, 4, 2, True),
(16, 3072, 1536): (16, 128, 3, 4, 4, 2, True),
(64, 3072, 1536): (16, 128, 3, 4, 4, 2, True),
}
_CACHE_FUSED_SHAPES = {
(8, 2112, 7168),
(16, 2112, 7168),
}
_oc = {}
def custom_kernel(data: input_t) -> output_t:
A = data[0]; Bq = data[2]; Bs = data[4]
m, k = A.shape; n = Bq.shape[0]
if not A.is_contiguous(): A = A.contiguous()
Bu = Bq.view(_u8); bsu = Bs.view(_u8); sn = bsu.shape[1]
pc = _PERSISTENT_CFGS.get((m, n, k)); lc = _LEAN_CFGS.get((m, n, k))
if pc is not None:
BM, BN, KI, nw, ns = pc
Af, Asr = dynamic_mxfp4_quant(A); Aq = Af.view(_u8); As = Asr.view(_u8)
ok = (m, n); out = _oc.get(ok)
if out is None: out = torch.empty((m, n), dtype=_bf16, device=A.device); _oc[ok] = out
tt = triton.cdiv(m, BM) * triton.cdiv(n, BN); nc = min(NUM_SMS, tt)
_persistent_cache_lean_kernel[(nc,)](
Aq, Bu, As, bsu, out, m, n, k, tt,
Aq.stride(0), Aq.stride(1), Bu.stride(1), Bu.stride(0),
As.stride(0), As.stride(1), out.stride(0), out.stride(1), sn,
BM=BM, BN=BN, BK=512, KI=KI, NSMS=nc, num_warps=nw, num_stages=ns, waves_per_eu=2)
elif lc is not None:
BM, BN, KI, gsm, nw, ns, even = lc
Af, Asr = dynamic_mxfp4_quant(A); Aq = Af.view(_u8); As = Asr.view(_u8)
ok = (m, n); out = _oc.get(ok)
if out is None: out = torch.empty((m, n), dtype=_bf16, device=A.device); _oc[ok] = out
grid = triton.cdiv(m, BM) * triton.cdiv(n, BN)
_lean_gemm_kernel[(grid,)](
Aq, Bu, As, bsu, out, m, n, k,
Aq.stride(0), Aq.stride(1), Bu.stride(1), Bu.stride(0),
As.stride(0), As.stride(1), out.stride(0), out.stride(1), sn,
BM=BM, BN=BN, BK=512, GSM=gsm, KI=KI, EVEN=even, num_warps=nw, num_stages=ns, waves_per_eu=2)
else:
BK = 512; fc = _FUSED_CFGS.get((m, n, k))
if fc: BM, BN, SK, KI, nw, ns, even = fc
else: BM, BN = 16, 128; SK = max(1, k // BK) if k >= 4096 else 1; KI = max(1, triton.cdiv(k, max(SK, 1) * BK)); nw, ns, even = 4, 1, False
grid = triton.cdiv(m, BM) * triton.cdiv(n, BN)
if SK <= 1:
ok = (m, n); out = _oc.get(ok)
if out is None: out = torch.empty((m, n), dtype=_bf16, device=A.device); _oc[ok] = out
_fused_gemm_kernel[(grid,)](
A, Bu, out, bsu, m, n, k, A.stride(0), A.stride(1), Bu.stride(1), Bu.stride(0), 0, out.stride(0), out.stride(1), sn, BLOCK_M=BM, BLOCK_N=BN, BLOCK_K=BK, GROUP_SIZE_M=4, NUM_KSPLIT=1, K_ITERS_PER_SPLIT=KI, EVEN_MNK=even, num_warps=nw, num_stages=ns, waves_per_eu=2)
else:
p = torch.empty((SK, m, n), dtype=_f32, device=A.device); out = torch.empty((m, n), dtype=_bf16, device=A.device)
if (m, n, k) in _CACHE_FUSED_SHAPES:
_fused_cache_gemm_kernel[(SK * grid,)](
A, Bu, p, bsu, m, n, k, A.stride(0), A.stride(1), Bu.stride(1), Bu.stride(0), p.stride(0), p.stride(1), p.stride(2), sn, BLOCK_M=BM, BLOCK_N=BN, BLOCK_K=BK, GROUP_SIZE_M=4, NUM_KSPLIT=SK, K_ITERS_PER_SPLIT=KI, EVEN_MNK=even, num_warps=nw, num_stages=ns, waves_per_eu=2)
else:
_fused_gemm_kernel[(SK * grid,)](
A, Bu, p, bsu, m, n, k, A.stride(0), A.stride(1), Bu.stride(1), Bu.stride(0), p.stride(0), p.stride(1), p.stride(2), sn, BLOCK_M=BM, BLOCK_N=BN, BLOCK_K=BK, GROUP_SIZE_M=4, NUM_KSPLIT=SK, K_ITERS_PER_SPLIT=KI, EVEN_MNK=even, num_warps=nw, num_stages=ns, waves_per_eu=2)
_reduce_kernel[(triton.cdiv(m, 16), triton.cdiv(n, BN))](p, out, m, n, p.stride(0), p.stride(1), p.stride(2), out.stride(0), out.stride(1), NS=SK, BM=16, BN=BN)
return out
scrolls · 287 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