Skip to content
KernelIndex
Search⌘K

submission 748551

Chivier · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:99e23ceff36841297414fc91d577e996422ee8bcad7539a4a2c01d23f1dfe18c
license declaredunknown
license concludedunknown
authorsChivier
imported2026-08-15

Techniques

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

fp4MXFP4 GEMM V276 — F32 CVT variant + int-ptr replay + minimal hot path.
tile-m = 16C.stride(0), C.stride(1), KS=ks, BM=16, BN=128,
tile-n = 64(64, 7168, 2048): (16, 128, 512, 4, 2, 8, 1), # V234: 12.0µs (BN=64 regressed to 13.5)

Kernel source

submission.py356 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
MXFP4 GEMM V276 — F32 CVT variant + int-ptr replay + minimal hot path.

CVT change: V_CVT_SCALEF32_PK_FP4_F32 (opcode 573) takes two f32 inputs directly.
Eliminates: bf16 conversion + uint16 bitcast + split + pack = ~6 ops per pair.
F32 variant: f32_to_fp4_scale(S0.f32, scale) — full f32 precision, no bf16 intermediate.
Also: pre-cache bs_sh view, minimal Python hot path.
"""
import os
os.environ['HIP_FORCE_DEV_KERNARG'] = '1'

import torch
import triton
import triton.language as tl
from task import input_t, output_t

# ─── Inlined _mxfp4_quant_op with bit-ops replacing log2/exp2 ───
@triton.jit
def _mxfp4_quant_op(
    x,
    BLOCK_SIZE_N,
    BLOCK_SIZE_M,
    MXFP4_QUANT_BLOCK_SIZE,
):
    """F32 CVT quant: software f32 scale + V_CVT_SCALEF32_PK_FP4_F32.

    F32 variant: takes two separate f32 values + f32 scale → packed FP4.
    No bf16 conversion needed — full f32 precision throughout.
    ISA opcode 573: f32_to_fp4_scale(S0.f32, scale) + f32_to_fp4_scale(S1.f32, scale)
    """
    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
    HQB: tl.constexpr = MXFP4_QUANT_BLOCK_SIZE // 2

    x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)

    # ── Step 1: amax + E8M0 scale ──
    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_exp = (amax >> 23) & 0xFF
    scale_e8m0_unbiased = amax_exp.to(tl.int32) - 129
    scale_e8m0_unbiased = tl.maximum(scale_e8m0_unbiased, -127)
    scale_e8m0_unbiased = tl.minimum(scale_e8m0_unbiased, 127)
    bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127

    # ── Step 2: Inverse scale ──
    inv_exp = (256 - amax_exp).to(tl.int32)
    inv_exp = tl.maximum(inv_exp, 0)
    inv_exp = tl.minimum(inv_exp, 254)
    quant_scale = (inv_exp.to(tl.uint32) << 23).to(tl.float32, bitcast=True)
    amax_f = amax.to(tl.float32, bitcast=True)
    quant_scale = tl.where(amax_f > 0, quant_scale, tl.zeros_like(quant_scale))

    # ── Step 3: Scale in f32 ──
    qx = x * quant_scale  # [BM, NQB, QBS] f32

    # ── Step 4: F32 CVT — skip bf16 conversion entirely ──
    # Split into even/odd f32 values for CVT_PK_FP4_F32(S0, S1, scale)
    qx_pairs = qx.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, HQB, 2)
    qx_even, qx_odd = tl.split(qx_pairs)  # each [BM, NQB, HQB, 1]
    qx_even = qx_even.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, HQB)
    qx_odd = qx_odd.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, HQB)
    one_f32 = tl.full([BLOCK_SIZE_M, NUM_QUANT_BLOCKS, HQB], 1.0, dtype=tl.float32)

    # V_CVT_SCALEF32_PK_FP4_F32: vdst, S0(f32), S1(f32), S2(scale_f32)
    # Output: byte at OPSEL position = {fp4(S1), fp4(S0)}
    fp4_u32 = tl.inline_asm_elementwise(
        "v_cvt_scalef32_pk_fp4_f32 $0, $1, $2, $3",
        "=v,v,v,v",
        args=[qx_even, qx_odd, one_f32],
        dtype=tl.uint32,
        is_pure=True,
        pack=1,
    )

    x_fp4 = (fp4_u32 & 0xFF).to(tl.uint8)
    x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
    return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)

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

@triton.jit
def _fused(
    A, Bq, Bs_sh, C,
    M, N, K,
    sa0, sa1, sbq0, sbq1, sc0, sc1,
    BM: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr,
    GM: tl.constexpr, NX: tl.constexpr,
    EVEN_K: tl.constexpr,
):
    pid = tl.program_id(0)
    nm = tl.cdiv(M, BM); nn = tl.cdiv(N, BN)
    pid = _remap_xcd(pid, nm * nn, NX)
    gn = GM * nn; gid = pid // gn; fm = gid * GM
    gsm = min(nm - fm, GM); pm = fm + (pid % gsm)
    pn = (pid % gn) // gsm
    if pm >= nm or pn >= nn: return
    om = pm * BM + tl.arange(0, BM)
    on = pn * BN + tl.arange(0, BN)
    acc = tl.zeros((BM, BN), dtype=tl.float32)
    HK: tl.constexpr = BK // 2
    SK: tl.constexpr = BK // 32
    sh_n_base = (on // 32) * K + (on % 16) * 4 + (on % 32) // 16
    for ks in range(0, K, BK):
        ok = ks + tl.arange(0, BK)
        if EVEN_K:
            a = tl.load(A + om[:, None] * sa0 + ok[None, :] * sa1)
        else:
            a = tl.load(A + om[:, None] * sa0 + ok[None, :] * sa1,
                        mask=(om[:, None] < M) & (ok[None, :] < K), other=0.0)
        aq, asc = _mxfp4_quant_op(a.to(tl.float32), BK, BM, 32)
        kp = ks // 2 + tl.arange(0, HK)
        if EVEN_K:
            b = tl.load(Bq + kp[:, None] * sbq1 + on[None, :] * sbq0,
                        cache_modifier=".cg")
        else:
            b = tl.load(Bq + kp[:, None] * sbq1 + on[None, :] * sbq0,
                        mask=on[None, :] < N, other=0,
                        cache_modifier=".cg")
        ksc = ks // 32 + tl.arange(0, SK)
        sh_k_part = (ksc // 8) * 256 + (ksc % 4) * 64 + ((ksc % 8) // 4) * 2
        if EVEN_K:
            bs = tl.load(Bs_sh + sh_n_base[:, None] + sh_k_part[None, :],
                         cache_modifier=".cg")
        else:
            bs = tl.load(Bs_sh + sh_n_base[:, None] + sh_k_part[None, :],
                         mask=on[:, None] < N, other=127,
                         cache_modifier=".cg")
        acc = tl.dot_scaled(aq, asc, "e2m1", b, bs, "e2m1", acc=acc, out_dtype=tl.float32)
    cp = C + om[:, None] * sc0 + on[None, :] * sc1
    if EVEN_K:
        tl.store(cp, acc.to(tl.bfloat16))
    else:
        tl.store(cp, acc.to(tl.bfloat16), mask=(om[:, None] < M) & (on[None, :] < N))

# ─── Split-K kernel for Case 2 (inherited from V39) ───
@triton.jit
def _fused_sk(
    A, Bq, Bs_sh, P,
    M, N, K,
    sa0, sa1, sbq0, sbq1, sp0, sp1, sp2,
    BM: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr,
    GM: tl.constexpr, NX: tl.constexpr, KS: tl.constexpr,
    EVEN_K: tl.constexpr,
):
    pid = tl.program_id(0)
    nm = tl.cdiv(M, BM); nn = tl.cdiv(N, BN)
    GRID = nm * nn * KS
    pid = _remap_xcd(pid, GRID, NX)
    pk = pid % KS; pmn = pid // KS
    pm = pmn // nn; pn = pmn % nn
    if pm >= nm or pn >= nn: return
    om = pm * BM + tl.arange(0, BM)
    on = pn * BN + tl.arange(0, BN)
    acc = tl.zeros((BM, BN), dtype=tl.float32)
    HK: tl.constexpr = BK // 2
    SK: tl.constexpr = BK // 32
    sh_n_base = (on // 32) * K + (on % 16) * 4 + (on % 32) // 16
    niters = K // (BK * KS)
    for ki in range(niters):
        ks = (ki * KS + pk) * BK
        ok = ks + tl.arange(0, BK)
        if EVEN_K:
            a = tl.load(A + om[:, None] * sa0 + ok[None, :] * sa1)
        else:
            a = tl.load(A + om[:, None] * sa0 + ok[None, :] * sa1,
                        mask=(om[:, None] < M) & (ok[None, :] < K), other=0.0)
        aq, asc = _mxfp4_quant_op(a.to(tl.float32), BK, BM, 32)
        kp = ks // 2 + tl.arange(0, HK)
        if EVEN_K:
            b = tl.load(Bq + kp[:, None] * sbq1 + on[None, :] * sbq0,
                        cache_modifier=".cg")
        else:
            b = tl.load(Bq + kp[:, None] * sbq1 + on[None, :] * sbq0,
                        mask=on[None, :] < N, other=0,
                        cache_modifier=".cg")
        ksc = ks // 32 + tl.arange(0, SK)
        sh_k_part = (ksc // 8) * 256 + (ksc % 4) * 64 + ((ksc % 8) // 4) * 2
        if EVEN_K:
            bs = tl.load(Bs_sh + sh_n_base[:, None] + sh_k_part[None, :],
                         cache_modifier=".cg")
        else:
            bs = tl.load(Bs_sh + sh_n_base[:, None] + sh_k_part[None, :],
                         mask=on[:, None] < N, other=127,
                         cache_modifier=".cg")
        acc = tl.dot_scaled(aq, asc, "e2m1", b, bs, "e2m1", acc=acc, out_dtype=tl.float32)
    pp = P + pk * sp0 + om[:, None] * sp1 + on[None, :] * sp2
    if EVEN_K:
        tl.store(pp, acc)
    else:
        tl.store(pp, acc, mask=(om[:, None] < M) & (on[None, :] < N))

@triton.jit
def _reduce(P, C, M, N, sp0, sp1, sp2, sc0, sc1,
            KS: tl.constexpr, BM: tl.constexpr, BN: 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)
    mask = (om[:, None] < M) & (on[None, :] < N)
    acc = tl.zeros((BM, BN), dtype=tl.float32)
    for s in range(KS):
        acc += tl.load(P + s * sp0 + om[:, None] * sp1 + on[None, :] * sp2, mask=mask, other=0.0)
    tl.store(C + om[:, None] * sc0 + on[None, :] * sc1, acc.to(tl.bfloat16), mask=mask)

# ─── Per-shape configs: (BM, BN, BK, nw, ns, GM, KS) ───
_CFG = {
    # K=512: BK=512 from V200 (single K-iter, proven 6.9-7.1µs)
    (4,   2880, 512):  (16, 16, 512, 4, 2, 1, 1),   # V200: 6.90µs
    (32,  4096, 512):  (16, 64, 512, 4, 2, 1, 1),   # V200: 7.13µs
    (32,  2880, 512):  (16, 64, 512, 4, 2, 1, 1),   # V200: 7.14µs
    # Large K: V278 best-of configs
    (16,  2112, 7168): (16, 64, 512, 4, 2, 8, 7),   # V277: 10.2µs!! (was V39: 14.1µs, -32%)
    (64,  7168, 2048): (16, 128, 512, 4, 2, 8, 1),  # V234: 12.0µs (BN=64 regressed to 13.5)
    (256, 3072, 1536): (16, 128, 256, 4, 3, 8, 1),  # V234: 14.0µs (BN=64 regressed to 15.0)
}

_out = {}
_bq = {}
_replay = {}   # (m,n,k) → (orig_launch, args_list, ptr_indices)
_ncall = {}

_bs_cache = {}

def _get_bq(bq_raw):
    k = id(bq_raw)
    c = _bq.get(k)
    if c is not None and c[0] is bq_raw: return c[1]
    r = bq_raw.view(torch.uint8)
    _bq[k] = (bq_raw, r)
    return r

def _get_bs(bs_raw):
    k = id(bs_raw)
    c = _bs_cache.get(k)
    if c is not None and c[0] is bs_raw: return c[1]
    r = bs_raw.view(torch.uint8)
    _bs_cache[k] = (bs_raw, r)
    return r

def _get_out(m, n, dev):
    k = (m, n)
    c = _out.get(k)
    if c is not None: return c
    c = torch.empty((m, n), dtype=torch.bfloat16, device=dev)
    _out[k] = c
    return c

def _triton_call(A, bq, bs_sh, C, m, n, k, BM, BN, BK, nw, ns, gm, even_k, grid):
    _fused[(grid,)](
        A, bq, bs_sh, C, m, n, k,
        A.stride(0), A.stride(1), bq.stride(0), bq.stride(1),
        C.stride(0), C.stride(1),
        BM=BM, BN=BN, BK=BK, GM=gm, NX=8,
        EVEN_K=even_k, num_warps=nw, num_stages=ns,
    )

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[0]
    bq = _get_bq(B_q)
    bs_sh = _get_bs(B_scale_sh)
    cfg = _CFG.get((m, n, k))
    if cfg is None:
        cfg = (16, 128, 256, 4, 3, 8, 1)
    BM, BN, BK, nw, ns, gm, ks = cfg
    even_k = (k % BK == 0) and (m % BM == 0) and (n % BN == 0)

    if ks > 1:
        grid_mn = triton.cdiv(m, BM) * triton.cdiv(n, BN)
        P = torch.empty((ks, m, n), device=A.device, dtype=torch.float32)
        _fused_sk[(grid_mn * ks,)](
            A, bq, bs_sh, P, m, n, k,
            A.stride(0), A.stride(1), bq.stride(0), bq.stride(1),
            P.stride(0), P.stride(1), P.stride(2),
            BM=BM, BN=BN, BK=BK, GM=gm, NX=8, KS=ks,
            EVEN_K=even_k, num_warps=nw, num_stages=ns,
        )
        C = _get_out(m, n, A.device)
        _reduce[(triton.cdiv(m, 16), triton.cdiv(n, 128))](
            P, C, m, n, P.stride(0), P.stride(1), P.stride(2),
            C.stride(0), C.stride(1), KS=ks, BM=16, BN=128,
        )
        return C

    C = _get_out(m, n, A.device)
    grid = triton.cdiv(m, BM) * triton.cdiv(n, BN)
    sk = (m, n, k)

    # ─── Fast replay: pass data_ptr() ints instead of Tensors ───
    # C launcher: PyLong_Check → PyLong_AsULL (0.01µs)
    # vs Tensor: getAttr("data_ptr") → Call → AsULL → hipPointerGetAttribute (0.5µs)
    rp = _replay.get(sk)
    if rp is not None:
        orig_fn, tmpl, pidx = rp
        a2 = list(tmpl)
        # Pass int ptrs (fast path in C launcher)
        a2[pidx[0]] = A.data_ptr()
        a2[pidx[1]] = bq.data_ptr()
        a2[pidx[2]] = bs_sh.data_ptr()
        a2[pidx[3]] = C.data_ptr()
        a2[1] = grid  # update gridX
        try:
            orig_fn(*a2)
            return C
        except Exception as e:
            import sys
            print(f"[V274] replay FAIL: {e}", file=sys.stderr)
            del _replay[sk]

    cnt = _ncall.get(sk, 0)
    _ncall[sk] = cnt + 1

    if cnt == 0:
        # 1st call: compile + install capture
        import sys
        a_ptr = A.data_ptr(); bq_ptr = bq.data_ptr()
        bs_ptr = bs_sh.data_ptr(); c_ptr = C.data_ptr()
        ret = _fused[(grid,)](
            A, bq, bs_sh, C, m, n, k,
            A.stride(0), A.stride(1), bq.stride(0), bq.stride(1),
            C.stride(0), C.stride(1),
            BM=BM, BN=BN, BK=BK, GM=gm, NX=8,
            EVEN_K=even_k, num_warps=nw, num_stages=ns,
        )
        if ret is not None and hasattr(ret, 'run') and hasattr(ret.run, 'launch'):
            orig = ret.run.launch
            _run_obj = ret.run  # hold reference to restore later
            def _cap(*args, _sk=sk, _orig=orig, _run=_run_obj):
                import sys as _sys, torch as _torch
                pidx = []
                for i, v in enumerate(args):
                    if isinstance(v, _torch.Tensor):
                        pidx.append(i)
                if len(pidx) >= 4:
                    _replay[_sk] = (_orig, list(args), pidx[:4])
                    # Restore original launch to stop re-capturing
                    _run.launch = _orig
                    _sys.stderr.write(f"[V274] CAPTURED {_sk}: pidx={pidx[:4]}\n")
                return _orig(*args)
            ret.run.launch = _cap
        return C

    # Normal dispatch
    _triton_call(A, bq, bs_sh, C, m, n, k, BM, BN, BK, nw, ns, gm, even_k, grid)
    return C
scrolls · 356 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