Skip to content
KernelIndex
Search⌘K

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
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
9.99µs
#233 of 1143
2026-04-04

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.

fp4MXFP4 GEMM - v58: selective cache modifiers on winning paths only.
persistent-kernel- persistent K=2048 path
split-k- fused split-K path for the K=7168 shapes
tile-k = 512BM=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 = 16else: 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

Kernel 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