Skip to content
KernelIndex
Search⌘K

submission 720499

Lemonade · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 239 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-720499?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
10.5µs
#268 of 1143
2026-04-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:2238c1fff5e923cea2085a91f91abd99c7c1222e18cc51b64e578a0745988a04
license declaredunknown
license concludedunknown
authorsLemonade
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp4MXFP4-MM v27: Maximum performance with pre-allocated buffers and minimal overhead.
split-kGROUP_SIZE_M: tl.constexpr, NUM_KSPLIT: tl.constexpr, SPLITK_SIZE: tl.constexpr,
stages = 2NUM_KSPLIT=1, SPLITK_SIZE=k, num_warps=nw, num_stages=2)
tile-k = 512- Fused quant+GEMM for m<=32 with BK=512, BN=128 for m<=16 large K
tile-m = 16BM = 16; nw = 2
tile-n = 128- Fused quant+GEMM for m<=32 with BK=512, BN=128 for m<=16 large K

Kernel source

submission.py239 lines
"""
MXFP4-MM v27: Maximum performance with pre-allocated buffers and minimal overhead.
- Pre-allocates ALL intermediate buffers on first call (zero alloc on hot path)
- Caches B reshapes across calls with same shape
- Fused quant+GEMM for m<=32 with BK=512, BN=128 for m<=16 large K
- CK GEMM for m>32 with pre-allocated output
- All view/reshape operations cached
"""
from task import input_t, output_t
import os
os.environ["CU_NUM"] = "256"
import torch
import triton
import triton.language as tl

# Global buffer cache - avoids torch.empty() on hot path
_cache = {}


@triton.jit
def _fused_qgemm(
    a_ptr, b_ptr, c_ptr, b_scale_ptr,
    M, N, K,
    stride_am, stride_ak, stride_bn, stride_bk,
    stride_ck, stride_cm, stride_cn, stride_bsn, stride_bsk,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
    GROUP_SIZE_M: tl.constexpr, NUM_KSPLIT: tl.constexpr, SPLITK_SIZE: tl.constexpr,
    NUM_XCDS: tl.constexpr = 8,
):
    SCALE_GROUP_SIZE: tl.constexpr = 32
    tl.assume(stride_am > 0); tl.assume(stride_ak > 0)
    tl.assume(stride_bn > 0); tl.assume(stride_bk > 0)
    tl.assume(stride_cm > 0); tl.assume(stride_cn > 0)
    tl.assume(stride_bsn > 0); tl.assume(stride_bsk > 0)
    pid_raw = tl.program_id(0)
    num_pid_m = tl.cdiv(M, BLOCK_M); num_pid_n = tl.cdiv(N, BLOCK_N)
    GRID_MN = num_pid_m * num_pid_n; GT = GRID_MN * NUM_KSPLIT
    ppx = (GT + NUM_XCDS - 1) // NUM_XCDS
    tx = GT % NUM_XCDS; tx = NUM_XCDS if tx == 0 else tx
    xcd = pid_raw % NUM_XCDS; lp = pid_raw // NUM_XCDS
    if xcd < tx: pu = xcd * ppx + lp
    else: pu = tx * ppx + (xcd - tx) * (ppx - 1) + lp
    pid_k = pu % NUM_KSPLIT; pid = pu // NUM_KSPLIT
    if NUM_KSPLIT == 1 and GROUP_SIZE_M > 1:
        npig = GROUP_SIZE_M * num_pid_n; gid = pid // npig
        fpm = gid * GROUP_SIZE_M; gsm = min(num_pid_m - fpm, GROUP_SIZE_M)
        tl.assume(gsm >= 0)
        pid_m = fpm + (pid % gsm); pid_n = (pid % npig) // gsm
    else:
        pid_m = pid // num_pid_n; pid_n = pid % num_pid_n
    tl.assume(pid_m >= 0); tl.assume(pid_n >= 0); tl.assume(pid_k >= 0)
    if (pid_k * SPLITK_SIZE // 2) < K:
        nki = tl.cdiv(SPLITK_SIZE // 2, BLOCK_K // 2)
        offs_m = (pid_m * BLOCK_M + tl.arange(0, BLOCK_M)) % M
        ok_bf = pid_k * SPLITK_SIZE + tl.arange(0, BLOCK_K)
        a_ptrs = a_ptr + offs_m[:, None] * stride_am + ok_bf[None, :] * stride_ak
        obn = (pid_n * (BLOCK_N // 16) + tl.arange(0, BLOCK_N // 16)) % N
        oks = pid_k * (SPLITK_SIZE // 2) * 16 + tl.arange(0, (BLOCK_K // 2) * 16)
        b_ptrs = b_ptr + obn[:, None] * stride_bn + oks[None, :] * stride_bk
        obsn = (pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)) % N
        obsk = pid_k * (SPLITK_SIZE // SCALE_GROUP_SIZE * 32) + tl.arange(0, BLOCK_K // SCALE_GROUP_SIZE * 32)
        bs_ptrs = b_scale_ptr + obsn[:, None] * stride_bsn + obsk[None, :] * stride_bsk
        acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
        DM: tl.constexpr = 149 << 23
        DF: tl.constexpr = tl.cast(DM, tl.float32, bitcast=True)
        for _ in range(pid_k * nki, (pid_k + 1) * nki):
            ab = tl.load(a_ptrs)
            af = ab.reshape(BLOCK_M * (BLOCK_K // 32), 32).to(tl.float32)
            ax = tl.max(tl.abs(af), axis=1, keep_dims=True)
            ax = ax.to(tl.int32, bitcast=True)
            ax = (ax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
            ax = ax.to(tl.float32, bitcast=True)
            su = tl.log2(ax).floor() - 2
            su = tl.clamp(su, min=-127, max=127)
            a_sc = (su.to(tl.uint8) + 127).reshape(BLOCK_M, BLOCK_K // 32)
            qx = af * tl.exp2(-su)
            qx = qx.to(tl.uint32, bitcast=True); sg = qx & 0x80000000; qx = qx ^ sg
            qf = qx.to(tl.float32, bitcast=True)
            st = qf >= 6; dn = (not st) & (qf < 1); nr = not (st | dn)
            dx = (qf + DF).to(tl.uint32, bitcast=True) - DM; dx = dx.to(tl.uint8)
            nx = qx.to(tl.int32, bitcast=True); mo = (nx >> 22) & 1
            nx += (-126 << 23) + (1 << 21) - 1; nx += mo; nx = (nx >> 22).to(tl.uint8)
            e = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)
            e = tl.where(nr, nx, e); e = tl.where(dn, dx, e)
            e = e | (sg >> 28).to(tl.uint8)
            e = tl.reshape(e, [BLOCK_M * (BLOCK_K // 32), 16, 2])
            ev, od = tl.split(e); afp4 = (ev | (od << 4)).reshape(BLOCK_M, BLOCK_K // 2)
            br = tl.load(b_ptrs)
            b = (br.reshape(1, BLOCK_N // 16, BLOCK_K // 64, 2, 16, 16)
                 .permute(0, 1, 4, 2, 3, 5).reshape(BLOCK_N, BLOCK_K // 2).trans(1, 0))
            bsc = (tl.load(bs_ptrs)
                   .reshape(BLOCK_N // 32, BLOCK_K // 32 // 8, 4, 16, 2, 2, 1)
                   .permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_N, BLOCK_K // 32))
            acc = tl.dot_scaled(afp4, a_sc, "e2m1", b, bsc, "e2m1", acc)
            a_ptrs += BLOCK_K * stride_ak
            b_ptrs += (BLOCK_K // 2) * 16 * stride_bk
            bs_ptrs += BLOCK_K * stride_bsk
        ocm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M).to(tl.int64)
        ocn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N).to(tl.int64)
        cm = (ocm[:, None] < M) & (ocn[None, :] < N)
        if NUM_KSPLIT == 1:
            tl.store(c_ptr + ocm[:, None] * stride_cm + ocn[None, :] * stride_cn,
                     acc.to(c_ptr.type.element_ty), mask=cm)
        else:
            tl.store(c_ptr + pid_k * stride_ck + ocm[:, None] * stride_cm + ocn[None, :] * stride_cn,
                     acc, mask=cm)


@triton.jit
def _reduce(cp, co, M, N, spk, spm, spn, som, son,
            BM: tl.constexpr, BN: tl.constexpr, NK: tl.constexpr, MK: tl.constexpr):
    pm = tl.program_id(0); pn = tl.program_id(1)
    om = (pm * BM + tl.arange(0, BM)) % M; on = (pn * BN + tl.arange(0, BN)) % N
    ok = tl.arange(0, MK)
    p = cp + ok[:, None, None] * spk + om[None, :, None] * spm + on[None, None, :] * spn
    v = tl.load(p, mask=ok[:, None, None] < NK) if NK != MK else tl.load(p)
    tl.store(co + om[:, None] * som + on[None, :] * son, tl.sum(v, axis=0).to(co.type.element_ty))


@triton.jit
def _qshuf(x_ptr, xf_ptr, bs_ptr, sxm, sxn, sfm, sfn, sbm, sbn,
           M: tl.constexpr, N: tl.constexpr, sN: tl.constexpr,
           sMP: tl.constexpr, sNP: tl.constexpr,
           BS: tl.constexpr, QBS: tl.constexpr):
    pm = tl.program_id(0); pn = tl.program_id(1)
    sxm = tl.cast(sxm, tl.int64); sxn = tl.cast(sxn, tl.int64)
    sfm = tl.cast(sfm, tl.int64); sfn = tl.cast(sfn, tl.int64)
    xm = pm * BS + tl.arange(0, BS); xn = pn * QBS + tl.arange(0, QBS)
    x = tl.load(x_ptr + xm[:, None] * sxm + xn[None, :] * sxn,
                mask=(xm < M)[:, None] & (xn < N)[None, :]).to(tl.float32)
    ax = tl.max(tl.abs(x), axis=1, keep_dims=True)
    ax = ax.to(tl.int32, bitcast=True)
    ax = (ax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
    ax = ax.to(tl.float32, bitcast=True)
    su = tl.log2(ax).floor() - 2; su = tl.clamp(su, min=-127, max=127)
    qx = x * tl.exp2(-su); bs = su.to(tl.uint8) + 127
    qx = qx.to(tl.uint32, bitcast=True); s = qx & 0x80000000; qx = qx ^ s
    qf = qx.to(tl.float32, bitcast=True)
    st = qf >= 6; dn = (not st) & (qf < 1); nr = not (st | dn)
    DM: tl.constexpr = 149 << 23; DF: tl.constexpr = tl.cast(DM, tl.float32, bitcast=True)
    dx = (qf + DF).to(tl.uint32, bitcast=True) - DM; dx = dx.to(tl.uint8)
    nx = qx.to(tl.int32, bitcast=True); mo = (nx >> 22) & 1
    nx += (-126 << 23) + (1 << 21) - 1; nx += mo; nx = (nx >> 22).to(tl.uint8)
    e = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)
    e = tl.where(nr, nx, e); e = tl.where(dn, dx, e); e = e | (s >> 28).to(tl.uint8)
    e = tl.reshape(e, [BS, QBS // 2, 2]); ev, od = tl.split(e); ot = ev | (od << 4)
    om = pm * BS + tl.arange(0, BS); on = pn * QBS // 2 + tl.arange(0, QBS // 2)
    tl.store(xf_ptr + om[:, None] * sfm + on[None, :] * sfn, ot,
             mask=(om < M)[:, None] & (on < (N // 2))[None, :])
    bm = pm * BS + tl.arange(0, BS); bn = pn
    b0 = bm[:, None] // 32; b12 = bm[:, None] % 32; b1 = b12 // 16; b2 = b12 % 16
    b3 = bn[None, :] // 8; b45 = bn[None, :] % 8; b4 = b45 // 4; b5 = b45 % 4
    bo = b1 + b4*2 + b2*4 + b5*64 + b3*256 + b0*32*sN
    m1 = (bm < M)[:, None] & (bn < sN)[None, :]
    m2 = (bm < sMP)[:, None] & (bn < sNP)[None, :]
    tl.store(bs_ptr + bo, tl.where(m1, bs, 127), mask=m2)


def _get_bufs(m, n, k, device):
    """Get pre-allocated buffers for a given shape. Avoids torch.empty() on hot path."""
    key = (m, n, k)
    if key not in _cache:
        from aiter import dtypes
        sM = triton.cdiv(m, 32) * 32
        sNv = triton.cdiv(k, 32)
        sN = triton.cdiv(sNv, 8) * 8
        _cache[key] = {
            'xf': torch.empty((m, k // 2), dtype=torch.uint8, device=device),
            'bs': torch.empty((triton.cdiv(m, 256) * 256, sN), dtype=torch.uint8, device=device),
            'C': torch.empty((m, n), dtype=torch.bfloat16, device=device),
            'sNv': sNv, 'sM': sM, 'sN': sN,
        }
    return _cache[key]


def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    m, k = A.shape; n, _ = B.shape

    # Cache B reshapes (same across calls)
    bkey = (id(B_shuffle), id(B_scale_sh))
    if bkey not in _cache:
        _cache[bkey] = (
            B_shuffle.view(torch.uint8).reshape(n // 16, (k // 2) * 16),
            B_scale_sh.view(torch.uint8).reshape(B_scale_sh.view(torch.uint8).shape[0] // 32,
                                                  B_scale_sh.view(torch.uint8).shape[1] * 32),
        )
    Bp, Bs = _cache[bkey]

    BN = 64

    if m <= 32:
        BM = 16; nw = 2
        BK = 512 if k >= 512 and k % 512 == 0 else 256
        if BK == 512: nw = 4
        if m <= 16 and k > 1024: BN = 128
        tiles = triton.cdiv(m, BM) * triton.cdiv(n, BN)
        NS = 1; mxk = k // (BK * 2)
        if mxk >= 2 and tiles < 128:
            while tiles * NS <= 256 and NS < mxk: NS *= 2
        NS = min(NS, max(1, min(mxk, 8)))
        SS = triton.cdiv(triton.cdiv(k, NS), BK) * BK
        AK = triton.cdiv(k, SS)
        grid = (AK * triton.cdiv(m, BM) * triton.cdiv(n, BN),)
        if AK == 1:
            C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
            _fused_qgemm[grid](A, Bp, C, Bs, m, n, k,
                A.stride(0), A.stride(1), Bp.stride(0), Bp.stride(1),
                0, C.stride(0), C.stride(1), Bs.stride(0), Bs.stride(1),
                BLOCK_M=BM, BLOCK_N=BN, BLOCK_K=BK, GROUP_SIZE_M=4,
                NUM_KSPLIT=1, SPLITK_SIZE=k, num_warps=nw, num_stages=2)
            return C
        else:
            Cp = torch.empty((AK, m, n), dtype=torch.float32, device=A.device)
            _fused_qgemm[grid](A, Bp, Cp, Bs, m, n, k,
                A.stride(0), A.stride(1), Bp.stride(0), Bp.stride(1),
                Cp.stride(0), Cp.stride(1), Cp.stride(2), Bs.stride(0), Bs.stride(1),
                BLOCK_M=BM, BLOCK_N=BN, BLOCK_K=BK, GROUP_SIZE_M=1,
                NUM_KSPLIT=AK, SPLITK_SIZE=SS, num_warps=nw, num_stages=2)
            C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
            MK = triton.next_power_of_2(AK)
            _reduce[(triton.cdiv(m, 16), triton.cdiv(n, 64))](
                Cp, C, m, n, Cp.stride(0), Cp.stride(1), Cp.stride(2),
                C.stride(0), C.stride(1), BM=16, BN=64, NK=AK, MK=MK)
            return C
    else:
        # Fast quant+shuffle + CK GEMM with pre-allocated buffers
        from aiter import dtypes
        import aiter
        bufs = _get_bufs(m, n, k, A.device)
        xf, bs_buf = bufs['xf'], bufs['bs']
        sNv, sM, sN = bufs['sNv'], bufs['sM'], bufs['sN']
        _qshuf[(triton.cdiv(m, 128), sNv)](
            A, xf, bs_buf, A.stride(0), A.stride(1), xf.stride(0), xf.stride(1),
            bs_buf.stride(0), bs_buf.stride(1), M=m, N=k, sN=sNv, sMP=sM, sNP=sN, BS=128, QBS=32)
        return aiter.gemm_a4w4(xf.view(dtypes.fp4x2), B_shuffle,
                               bs_buf.view(dtypes.fp8_e8m0), B_scale_sh,
                               dtype=dtypes.bf16, bpreshuffle=True)
scrolls · 239 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