Skip to content
KernelIndex
Search⌘K

submission 677752

xg · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-677752?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
181.6µs
#1137 of 1143
2026-03-31

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:1510080b7dc538b6e961927b26195fe55d70c1f0b88acfcba838679ec96a64d4
license declaredunknown
license concludedunknown
authorsxg
imported2026-08-26

Techniques

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

fp4Fused MXFP4 quant + GEMM: bf16 A -> inline MXFP4 quant -> tl.dot_scaled GEMM -> bf16 C.
tile-n = 16RBM, RBN = 16, 64

Kernel source

submission.py348 lines
"""
Fused MXFP4 quant + GEMM: bf16 A -> inline MXFP4 quant -> tl.dot_scaled GEMM -> bf16 C.
Uses only A, B_q, B_shuffle, B_scale_sh (plus B for API); no input cache, no host-side copies.
B layout matches shuffle_weight(16,16) + e8m0_shuffle (aiter.gemm_a4w4).
"""
from task import input_t, output_t
import torch
import triton
import triton.language as tl

SCALE_GROUP_SIZE = 32


@triton.jit
def _b_shuffle_phys_flat(n_l, k_l, KP):
    """Row-major [N, KP] byte index after shuffle_weight(..., layout=(16,16))."""
    k_blk = KP // 32
    nb = n_l // 16
    ni = n_l % 16
    kb = k_l // 32
    rem = k_l % 32
    sub = rem // 16
    ki = rem % 16
    return (((nb * k_blk + kb) * 2 + sub) * 16 + ni) * 16 + ki


@triton.jit
def _scale_shuffle_phys_flat(ml, nl, SN):
    """Linear index into row-major padded tensor from e8m0_shuffle (aiter fp4_utils)."""
    s1 = SN // 8
    a = ml // 32
    rem = ml % 32
    b = rem // 16
    c = rem % 16
    d = nl // 8
    rem2 = nl % 8
    e = rem2 // 4
    f = rem2 % 4
    return (((a * s1 + d) * 4 + f) * 16 + c) * 4 + e * 2 + b


@triton.jit
def _mxfp4_quant_inline(x, BSK: tl.constexpr, BSM: tl.constexpr, SGS: tl.constexpr):
    NQB: tl.constexpr = BSK // SGS
    x = x.reshape(BSM, NQB, SGS)
    amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
    amax = amax.to(tl.int32, bitcast=True)
    amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
    amax = amax.to(tl.float32, bitcast=True)
    su = tl.log2(amax).floor() - 2
    su = tl.clamp(su, min=-127, max=127)
    bs = su.to(tl.uint8) + 127
    qx = x * tl.exp2(-su)
    qx_u32 = qx.to(tl.uint32, bitcast=True)
    s = qx_u32 & 0x80000000
    qx_u32 = qx_u32 ^ s
    qf = qx_u32.to(tl.float32, bitcast=True)
    sat = qf >= 6
    den = (not sat) & (qf < 1)
    nor = not (sat | den)
    de: tl.constexpr = ((127 - 1) + (23 - 1) + 1) << 23
    df: tl.constexpr = tl.cast(de, tl.float32, bitcast=True)
    dx = qf + df
    dx = dx.to(tl.int32, bitcast=True)
    dx -= de
    dx = dx.to(tl.uint8)
    nx = qx_u32.to(tl.int32, bitcast=True)
    mo = (nx >> 22) & 1
    val_add: tl.constexpr = ((1 - 127) << 23) + (1 << 21) - 1
    nx += val_add
    nx += mo
    nx = nx >> 22
    nx = nx.to(tl.uint8)
    v = tl.full(qx_u32.type.get_block_shapes(), 0x7, dtype=tl.uint8)
    v = tl.where(nor, nx, v)
    v = tl.where(den, dx, v)
    sl = s >> 28
    sl = sl.to(tl.uint8)
    v = v | sl
    v = tl.reshape(v, [BSM, NQB, SGS // 2, 2])
    ev, od = tl.split(v)
    fp4 = ev | (od << 4)
    fp4 = fp4.reshape(BSM, BSK // 2)
    return fp4, bs.reshape(BSM, NQB)


@triton.jit
def _remap_xcd(pid, GM, NX: tl.constexpr = 8):
    ppx = (GM + NX - 1) // NX
    tx = GM % NX
    tx = NX if tx == 0 else tx
    xcd = pid % NX
    lp = pid // NX
    if xcd < tx:
        pid = xcd * ppx + lp
    else:
        pid = tx * ppx + (xcd - tx) * (ppx - 1) + lp
    return pid


@triton.jit
def _pgrid(pid, npm, npn, GSM: tl.constexpr = 1):
    if GSM == 1:
        pm = pid // npn
        pn = pid % npn
    else:
        npig = GSM * npn
        gi = pid // npig
        fpm = gi * GSM
        gsm = min(npm - fpm, GSM)
        tl.assume(gsm >= 0)
        pm = fpm + (pid % gsm)
        pn = (pid % npig) // gsm
    return pm, pn


@triton.heuristics({
    "EVEN_K": lambda a: (a["KP"] % (a["BSK"] // 2) == 0)
    and (a["SPBS"] % a["BSK"] == 0) and (a["KP"] % (a["SPBS"] // 2) == 0),
})
@triton.jit
def _fqg_kernel(
    a_ptr, b_ptr, c_ptr, bs_ptr,
    M, N, KP, SN, NUM_SG,
    sa0, sa1, sb0, sb1, sbs0, sbs1, sck, scm, scn,
    BSM: tl.constexpr, BSN: tl.constexpr, BSK: tl.constexpr,
    GSM: tl.constexpr, NKS: tl.constexpr, SPBS: tl.constexpr,
    EVEN_K: tl.constexpr,
    num_warps: tl.constexpr, num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
):
    tl.assume(sa0 > 0)
    tl.assume(sa1 > 0)
    tl.assume(sb0 > 0)
    tl.assume(sb1 > 0)
    tl.assume(sbs0 > 0)
    tl.assume(sbs1 > 0)
    tl.assume(scm > 0)
    tl.assume(scn > 0)
    GMN = tl.cdiv(M, BSM) * tl.cdiv(N, BSN)
    SGS: tl.constexpr = 32
    pu = tl.program_id(0)
    pu = _remap_xcd(pu, GMN * NKS, NX=8)
    pk = pu % NKS
    p = pu // NKS
    npm = tl.cdiv(M, BSM)
    npn = tl.cdiv(N, BSN)
    if NKS == 1:
        pm, pn = _pgrid(p, npm, npn, GSM=GSM)
    else:
        pm = p // npn
        pn = p % npn
    tl.assume(pm >= 0)
    tl.assume(pn >= 0)
    tl.assume(pk >= 0)
    if (pk * SPBS // 2) < KP:
        nki = tl.cdiv(SPBS // 2, BSK // 2)
        om = (pm * BSM + tl.arange(0, BSM)) % M
        on = (pn * BSN + tl.arange(0, BSN)) % N
        okp = tl.arange(0, BSK // 2)
        okb = tl.arange(0, BSK)
        oksb = pk * SPBS + okb
        ap = a_ptr + om[:, None] * sa0 + oksb[None, :] * sa1
        acc = tl.zeros((BSM, BSN), dtype=tl.float32)
        for ki in range(pk * nki, (pk + 1) * nki):
            inner = ki - pk * nki
            oksp = pk * (SPBS // 2) + inner * (BSK // 2) + okp
            oks = pk * (SPBS // SGS) + inner * (BSK // SGS) + tl.arange(0, BSK // SGS)
            n_l = on[None, :].to(tl.int64)
            k_l = oksp[:, None].to(tl.int64)
            phys_b = _b_shuffle_phys_flat(n_l, k_l, KP)
            bn = phys_b // KP
            bk = phys_b % KP
            bp = b_ptr + bn * sb0 + bk * sb1
            ml = on[:, None].to(tl.int64)
            nl = oks[None, :].to(tl.int64)
            phys_s = _scale_shuffle_phys_flat(ml, nl, SN)
            sr = phys_s // SN
            sc = phys_s % SN
            bsp = bs_ptr + sr * sbs0 + sc * sbs1
            mask_sc = oks[None, :] < NUM_SG
            if EVEN_K:
                ab = tl.load(ap).to(tl.float32)
            else:
                ab = tl.load(ap, mask=okb[None, :] < (KP * 2) - ki * BSK, other=0.0).to(tl.float32)
            af, asc = _mxfp4_quant_inline(ab, BSK, BSM, SGS)
            bsc = tl.load(bsp, mask=mask_sc, other=127)
            if EVEN_K:
                bv = tl.load(bp)
            else:
                bv = tl.load(bp, mask=okp[:, None] < KP - ki * (BSK // 2), other=0)
            acc = tl.dot_scaled(af, asc, "e2m1", bv, bsc, "e2m1", acc)
            ap += BSK * sa1
        c = acc.to(c_ptr.type.element_ty)
        ocm = pm * BSM + tl.arange(0, BSM).to(tl.int64)
        ocn = pn * BSN + tl.arange(0, BSN).to(tl.int64)
        cp = c_ptr + scm * ocm[:, None] + scn * ocn[None, :] + pk * sck
        cm = (ocm[:, None] < M) & (ocn[None, :] < N)
        tl.store(cp, c, mask=cm)


@triton.jit
def _red_kernel(
    yp, yo, M, N,
    syk, sym, syn, som, son,
    BM: tl.constexpr, BN: tl.constexpr,
    NKS: tl.constexpr, NKP: tl.constexpr,
):
    pm = tl.program_id(0)
    pn = tl.program_id(1)
    om = pm * BM + tl.arange(0, BM)
    on = pn * BN + tl.arange(0, BN)
    mk = (om[:, None] < M) & (on[None, :] < N)
    a = tl.zeros((BM, BN), dtype=tl.float32)
    for k in range(NKS):
        v = tl.load(yp + k * syk + om[:, None] * sym + on[None, :] * syn, mask=mk, other=0.0)
        a += v
    tl.store(yo + om[:, None] * som + on[None, :] * son, a.to(yo.type.element_ty), mask=mk)


def _get_spk(M, N, KP, BSK):
    CU = 304
    bm = 16 if M <= 16 else (32 if M <= 32 else (64 if M <= 64 else 128))
    t = ((M + bm - 1) // bm) * ((N + 127) // 128)
    c = CU / max(t, 1)
    s = 0
    while c >= pow(2, s + 1) and (pow(2, s + 1) * BSK) < 2 * KP:
        s += 1
    return min(s, 3)


def _spk_bs(KP, BSK, NKS):
    if NKS <= 1:
        return 2 * KP, BSK, 1
    SP = triton.cdiv((2 * triton.cdiv(KP, NKS)), BSK) * BSK
    b, n = BSK, NKS
    while n > 1 and b > 16:
        if KP % (SP // 2) == 0 and SP % b == 0 and KP % (b // 2) == 0:
            break
        elif KP % (SP // 2) != 0 and n > 1:
            n //= 2
        elif SP % b != 0:
            if n > 1:
                n //= 2
            elif b > 16:
                b //= 2
        elif KP % (b // 2) != 0 and b > 16:
            b //= 2
        else:
            break
        SP = triton.cdiv((2 * triton.cdiv(KP, n)), b) * b
    n = triton.cdiv(KP, (SP // 2))
    return SP, b, n


def _run_separate_path(A, B_q, B_scale_sh, B_shuffle, m, n, k):
    """For large M, use separate quant + aiter ASM GEMM (faster than fused for compute-bound cases)."""
    import aiter
    from aiter import dtypes
    from aiter.ops.triton.quant import dynamic_mxfp4_quant
    from aiter.utility.fp4_utils import e8m0_shuffle

    A_fp4, A_scale = dynamic_mxfp4_quant(A)
    A_scale_sh = e8m0_shuffle(A_scale)
    A_q = A_fp4.view(dtypes.fp4x2)
    A_sc = A_scale_sh.view(dtypes.fp8_e8m0)
    return aiter.gemm_a4w4(A_q, B_shuffle, A_sc, B_scale_sh, dtype=dtypes.bf16, bpreshuffle=True)


def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    m, k = A.shape
    n = B_q.shape[0]
    KP = k // 2

    if m >= 128:
        return _run_separate_path(A, B_q, B_scale_sh, B_shuffle, m, n, k)

    B_u8 = B_shuffle.view(torch.uint8)
    BS_u8 = B_scale_sh.view(torch.uint8)
    sb0, sb1 = B_u8.stride()
    sbs0, sbs1 = BS_u8.stride()
    sn = B_scale_sh.shape[1]
    num_sg = k // SCALE_GROUP_SIZE

    if m <= 16:
        BSM, BSN, BSK = 16, 128, 256
        GSM, nw, ns, wpe, mid = 1, 4, 2, 3, 16
        NKS = 2 ** _get_spk(m, n, KP, BSK)
    elif m <= 32:
        BSM, BSN, BSK = 32, 128, 256
        GSM, nw, ns, wpe, mid = 1, 4, 2, 3, 16
        NKS = 1
    elif m <= 64:
        BSM, BSN, BSK = 64, 256, 256
        GSM, nw, ns, wpe, mid = 1, 4, 3, 2, 32
        NKS = 1
    else:
        BSM, BSN, BSK = 128, 256, 256
        GSM, nw, ns, wpe, mid = 2, 4, 3, 2, 32
        NKS = 1

    BSK = max(BSK, 128)
    if BSK >= 2 * KP:
        BSK = triton.next_power_of_2(2 * KP)
        NKS = 1

    if NKS > 1:
        SPBS, BSK, NKS = _spk_bs(KP, BSK, NKS)
    else:
        SPBS = 2 * KP

    if NKS > 1:
        ypp = torch.empty((NKS, m, n), dtype=torch.float32, device=A.device)
        out = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
        tgt = ypp
    else:
        out = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
        ypp = None
        tgt = out

    grid = lambda META: (
        META['NKS'] * triton.cdiv(m, META['BSM']) * triton.cdiv(n, META['BSN']),
    )

    _fqg_kernel[grid](
        A, B_u8, tgt, BS_u8,
        m, n, KP, sn, num_sg,
        A.stride(0), A.stride(1), sb0, sb1, sbs0, sbs1,
        0 if NKS == 1 else ypp.stride(0),
        tgt.stride(-2) if NKS <= 1 else ypp.stride(1),
        tgt.stride(-1) if NKS <= 1 else ypp.stride(2),
        BSM=BSM, BSN=BSN, BSK=BSK,
        GSM=GSM, NKS=NKS, SPBS=SPBS,
        num_warps=nw, num_stages=ns, waves_per_eu=wpe, matrix_instr_nonkdim=mid,
    )

    if NKS > 1:
        RBM, RBN = 16, 64
        _red_kernel[(triton.cdiv(m, RBM), triton.cdiv(n, RBN))](
            ypp, out, m, n,
            ypp.stride(0), ypp.stride(1), ypp.stride(2),
            out.stride(0), out.stride(1),
            BM=RBM, BN=RBN, NKS=NKS, NKP=triton.next_power_of_2(NKS),
        )

    return out
scrolls · 348 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