Skip to content
KernelIndex
Search⌘K

submission 711568

vuxml · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v10d_hybrid.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-711568?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.96µs
#227 of 1143
2026-04-03

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:59c6f64defae9d216e84a9e6be65a0cc63cc0017673c7e2535c27aa140b1d2c1
license declaredunknown
license concludedunknown
authorsvuxml
imported2026-08-15

Techniques

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

fp4v10d: hybrid fused vs {quant-kernel + fp4-GEMM}. Autotune (L2-cold)
num-warps = 4BM=QBM, BK=QBK, num_warps=4)
split-kSPLIT_K: tl.constexpr, EVEN_N: tl.constexpr, PREQUANT: tl.constexpr,
tile-k = 16QBM, QBK = 16, min(256, k)

Kernel source

submission_v10d_hybrid.py326 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
v10d: hybrid fused vs {quant-kernel + fp4-GEMM}. Autotune (L2-cold)
picks per shape. All launches via HIPLauncher (skip JITFunction.run).

Fused is optimal for small m·n (quant cost amortized). Prequant is
optimal for large m·n (A quantized once, not per N-tile).
"""
import os, sys
os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
import warnings; warnings.filterwarnings("ignore")
import torch
import triton
import triton.language as tl

import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op

_L = lambda *a: print(*a, file=sys.stderr, flush=True)
_DRV = triton.runtime.driver.active
_GCS = getattr(_DRV, "get_current_" + chr(115) + "tream")


# e8m0_shuffle forward flat idx: view(sm//32,2,16, sn//8,2,4).permute(0,3,5,2,4,1)
@triton.jit
def _sh_row(r, sn8):
    return (r // 32) * (sn8 * 256) + (r % 16) * 4 + (r // 16) % 2


@triton.jit
def _sh_col(c):
    return (c // 8) * 256 + (c % 4) * 64 + (c // 4) % 2 * 2


@triton.jit
def _gemm_k(
    A, Asc, Bq, Bsc, C,
    M, N, K, sA_m, sAsc_m, sBq_n, sC_k, sC_m, sn8,
    BM: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr,
    SPLIT_K: tl.constexpr, EVEN_N: tl.constexpr, PREQUANT: tl.constexpr,
):
    pid = tl.program_id(0)
    num_n = tl.cdiv(N, BN)
    num_mn = tl.cdiv(M, BM) * num_n
    pid_k = pid // num_mn
    pid_mn = pid % num_mn
    pid_m = pid_mn // num_n
    pid_n = pid_mn % num_n

    offs_m = pid_m * BM + tl.arange(0, BM)
    offs_n = pid_n * BN + tl.arange(0, BN)
    offs_n64 = offs_n.to(tl.int64)
    mask_m = offs_m < M
    mask_n = offs_n < N
    rk = tl.arange(0, BK)
    rk2 = tl.arange(0, BK // 2)
    rk32 = tl.arange(0, BK // 32)

    k_per = tl.cdiv(tl.cdiv(K, BK), SPLIT_K) * BK
    k_lo = pid_k * k_per
    k_hi = min(k_lo + k_per, K)

    bq_ptrs = Bq + offs_n64[:, None] * sBq_n + (k_lo // 2 + rk2)[None, :]
    bsc_row = _sh_row(offs_n64, sn8)
    acc = tl.zeros((BM, BN), dtype=tl.float32)

    if PREQUANT:
        a_ptrs = A + offs_m[:, None].to(tl.int64) * sA_m + (k_lo // 2 + rk2)[None, :]
        asc_ptrs = Asc + offs_m[:, None].to(tl.int64) * sAsc_m + (k_lo // 32 + rk32)[None, :]
    else:
        a_ptrs = A + offs_m[:, None].to(tl.int64) * sA_m + (k_lo + rk)[None, :]

    for k in tl.range(k_lo, k_hi, BK):
        if PREQUANT:
            a_fp4 = tl.load(a_ptrs, mask=mask_m[:, None], other=0)
            a_sc = tl.load(asc_ptrs, mask=mask_m[:, None], other=0)
        else:
            a_bf = tl.load(a_ptrs, mask=mask_m[:, None], other=0.0)
            a_fp4, a_sc = _mxfp4_quant_op(a_bf.to(tl.float32), BK, BM, 32)

        if EVEN_N:
            b_fp4_t = tl.load(bq_ptrs)
        else:
            b_fp4_t = tl.load(bq_ptrs, mask=mask_n[:, None], other=0)
        bsc_col = _sh_col(k // 32 + rk32)
        if EVEN_N:
            b_sc = tl.load(Bsc + bsc_row[:, None] + bsc_col[None, :])
        else:
            b_sc = tl.load(Bsc + bsc_row[:, None] + bsc_col[None, :],
                           mask=mask_n[:, None], other=0)

        acc = tl.dot_scaled(a_fp4, a_sc, "e2m1",
                            tl.trans(b_fp4_t), b_sc, "e2m1", acc)
        if PREQUANT:
            a_ptrs += BK // 2
            asc_ptrs += BK // 32
        else:
            a_ptrs += BK
        bq_ptrs += BK // 2

    c_off = (pid_k * sC_k + offs_m[:, None].to(tl.int64) * sC_m
             + offs_n[None, :])
    cmask = mask_m[:, None] & mask_n[None, :]
    if SPLIT_K == 1:
        tl.store(C + c_off, acc.to(tl.bfloat16), mask=cmask)
    else:
        tl.store(C + c_off, acc, mask=cmask)


@triton.jit
def _reduce_k(W, C, SK, M, N, sW_k, sW_m, sC_m,
              BLK: tl.constexpr, SKC: tl.constexpr):
    pid = tl.program_id(0)
    off = pid * BLK + tl.arange(0, BLK)
    om = off // N; on = off % N; mask = om < M
    base = om.to(tl.int64) * sW_m + on
    s = tl.zeros((BLK,), dtype=tl.float32)
    for i in tl.static_range(SKC):
        s += tl.load(W + i * sW_k + base, mask=mask & (i < SK), other=0.0)
    tl.store(C + om.to(tl.int64) * sC_m + on, s.to(tl.bfloat16), mask=mask)


@triton.jit
def _quant_a_k(A, Afp4, Asc, M, K, sA_m, sAf_m, sAs_m,
               BM: tl.constexpr, BK: tl.constexpr):
    pid = tl.program_id(0)
    nk = tl.cdiv(K, BK)
    pm = pid // nk; pk = pid % nk
    offs_m = pm * BM + tl.arange(0, BM)
    offs_k = pk * BK + tl.arange(0, BK)
    mask_m = offs_m < M
    a = tl.load(A + offs_m[:, None].to(tl.int64) * sA_m + offs_k[None, :],
                mask=mask_m[:, None], other=0.0).to(tl.float32)
    af, asc = _mxfp4_quant_op(a, BK, BM, 32)
    tl.store(Afp4 + offs_m[:, None].to(tl.int64) * sAf_m
             + (pk * (BK // 2) + tl.arange(0, BK // 2))[None, :],
             af, mask=mask_m[:, None])
    tl.store(Asc + offs_m[:, None].to(tl.int64) * sAs_m
             + (pk * (BK // 32) + tl.arange(0, BK // 32))[None, :],
             asc, mask=mask_m[:, None])


def _cfgs(m, n, k):
    """(BM, BN, BK, SPLIT_K, nw, nK, ns, PREQUANT).
    m≤32: fused only. m≥64: prequant only, BM∈{32,64}."""
    out = []
    PQs = (False,) if m <= 32 else (True,) if m >= 128 else (False, True)
    for PQ in PQs:
        BMs = (16,) if not PQ else tuple(b for b in (32, 64, 128) if b <= m)
        for BM in BMs:
            for BN in (32, 64, 128, 256):
                if BN > n: continue
                bt = -(-m // BM) * -(-n // BN)
                for BK in (256, 512):
                    if BK > k: continue
                    nkit = k // BK
                    SKs = [1]
                    if bt < 200 and nkit >= 2:
                        tgt = max(1, 256 // bt)
                        for s in (2, 4, 8):
                            if s <= nkit and s <= tgt * 2: SKs.append(s)
                    for SK in SKs:
                        for nw in (4, 8):
                            for nK in ((16,) if BM == 16 else (16, 32)):
                                out.append((BM, BN, BK, SK, nw, nK, 2, PQ))
    seen, r = set(), []
    for c in out:
        if c not in seen: seen.add(c); r.append(c)
    return r


_L2FLUSH = torch.empty(512 * 1024 * 1024, dtype=torch.int8, device="cuda")


def _gpu_time_cold(fn, n_iter=6):
    for _ in range(2): fn()
    torch.cuda.synchronize()
    evs = [(torch.cuda.Event(True), torch.cuda.Event(True)) for _ in range(n_iter)]
    for e0, e1 in evs:
        _L2FLUSH.zero_()
        e0.record(); fn(); e1.record()
    torch.cuda.synchronize()
    return sum(e0.elapsed_time(e1) for e0, e1 in evs) * 1000.0 / n_iter


def _ref(A, B_shuffle, B_scale_sh):
    Aq, As = dynamic_mxfp4_quant(A)
    As = e8m0_shuffle(As)
    return aiter.gemm_a4w4(Aq.view(dtypes.fp4x2), B_shuffle,
                           As.view(dtypes.fp8_e8m0), B_scale_sh,
                           dtype=dtypes.bf16, bpreshuffle=True)


_STATE: dict = {}


def _mk_launcher(kernel, grid_x, const_args, **kw):
    ck = kernel.warmup(*const_args, grid=(grid_x,), **kw)
    ck._init_handles()
    return ck.run, ck.function, ck.packed_metadata


def _build(data):
    A, B, B_q, B_shuffle, B_scale_sh = data
    m, k = A.shape; n = B.shape[0]
    sn = B_scale_sh.shape[1]; sn8 = sn // 8
    dev = A.device

    Bq = B_q.contiguous().view(torch.uint8)
    Bsc = B_scale_sh.contiguous().view(torch.uint8).reshape(-1)
    C_bf = torch.empty((m, n), dtype=torch.bfloat16, device=dev)
    W = torch.empty((8, m, n), dtype=torch.float32, device=dev)
    Afp4 = torch.empty((m, k // 2), dtype=torch.uint8, device=dev)
    Asc = torch.empty((m, k // 32), dtype=torch.uint8, device=dev)

    ref_f = _ref(A, B_shuffle, B_scale_sh).float()
    mag = ref_f.abs().mean().item() + 1e-9

    try:
        lq = _GCS(None)
    except Exception:
        lq = _GCS(torch.cuda.current_device())

    QBM, QBK = 16, min(256, k)
    q_gx = triton.cdiv(m, QBM) * (k // QBK)
    q_run, q_fun, q_pmeta = _mk_launcher(
        _quant_a_k, q_gx,
        (A, Afp4, Asc, m, k, k, k // 2, k // 32),
        BM=QBM, BK=QBK, num_warps=4)

    r_gx = triton.cdiv(m * n, 256)
    r_run, r_fun, r_pmeta = _mk_launcher(
        _reduce_k, r_gx,
        (W, C_bf, 8, m, n, m * n, n, n),
        BLK=256, SKC=8, num_warps=4)

    def _do_quant(A_in):
        q_run(q_gx, 1, 1, lq, q_fun, q_pmeta, None, None, None,
              A_in, Afp4, Asc, m, k, k, k // 2, k // 32, QBM, QBK)

    def _do_reduce(SK):
        r_run(r_gx, 1, 1, lq, r_fun, r_pmeta, None, None, None,
              W, C_bf, SK, m, n, m * n, n, n, 256, 8)

    cfgs = _cfgs(m, n, k)
    _L(f"\n[v10d m={m} n={n} k={k}] {len(cfgs)} cfgs")

    best, best_t, best_go = None, float("inf"), None
    for cfg in cfgs:
        BM, BN, BK, SK, nw, nK, ns, PQ = cfg
        C_out = W if SK > 1 else C_bf
        sC_k = W.stride(0) if SK > 1 else 0
        sC_m = W.stride(1) if SK > 1 else n
        sAm = (k // 2) if PQ else k
        even_n = (n % BN == 0)
        gx = triton.cdiv(m, BM) * triton.cdiv(n, BN) * SK
        A_in = Afp4 if PQ else A
        try:
            run, fun, pmeta = _mk_launcher(
                _gemm_k, gx,
                (A_in, Asc, Bq, Bsc, C_out, m, n, k,
                 sAm, k // 32, k // 2, sC_k, sC_m, sn8),
                BM=BM, BN=BN, BK=BK, SPLIT_K=SK, EVEN_N=even_n,
                PREQUANT=PQ, num_warps=nw, num_stages=ns,
                matrix_instr_nonkdim=nK, waves_per_eu=0)
        except Exception as e:
            if best is None:
                _L(f"  {cfg}: COMPILE {type(e).__name__}: {str(e)[:100]}")
            continue

        cargs = (BM, BN, BK, SK, even_n, PQ)

        def _go(_A=A, _Bq=Bq, _Bsc=Bsc, _run=run, _fun=fun,
                _pmeta=pmeta, _cargs=cargs, _gx=gx, _PQ=PQ, _SK=SK,
                _C=C_out, _sAm=sAm, _sCk=sC_k, _sCm=sC_m):
            if _PQ:
                _do_quant(_A)
                _run(_gx, 1, 1, lq, _fun, _pmeta, None, None, None,
                     Afp4, Asc, _Bq, _Bsc, _C, m, n, k,
                     _sAm, k // 32, k // 2, _sCk, _sCm, sn8, *_cargs)
            else:
                _run(_gx, 1, 1, lq, _fun, _pmeta, None, None, None,
                     _A, Asc, _Bq, _Bsc, _C, m, n, k,
                     _sAm, k // 32, k // 2, _sCk, _sCm, sn8, *_cargs)
            if _SK > 1:
                _do_reduce(_SK)
            return C_bf

        try:
            out = _go()
            torch.cuda.synchronize()
            err = ((out.float() - ref_f).abs().mean() / mag).item()
            if err > 5e-3:
                if best is None: _L(f"  {cfg}: ERR {err:.2%}")
                continue
            t = _gpu_time_cold(_go)
            if t < best_t:
                best_t, best, best_go = t, cfg, _go
                _L(f"  {cfg}: {t:.2f}us grid={gx} *")
        except Exception as e:
            if best is None:
                _L(f"  {cfg}: RUN {type(e).__name__}: {str(e)[:100]}")

    if best is None:
        _L(f"  → fallback")
        return {"hot": None}

    _L(f"  → best={best} @ {best_t:.2f}us")
    return {"hot": best_go}


def custom_kernel(data):
    A = data[0]; Bq = data[2]
    key = (A.shape[0], Bq.shape[0], A.shape[1])
    S = _STATE.get(key)
    if S is None:
        S = _build(data); _STATE[key] = S
    hot = S["hot"]
    if hot is None:
        return _ref(A, data[3], data[4])
    return hot(A, Bq, data[4])
scrolls · 326 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 705642.

⋯ 151 unchanged lines
for PQ in PQs:
BMs = (16,) if not PQ else tuple(b for b in (32, 64, 128) if b <= m)
for BM in BMs:
- for BN in (64, 128, 256):
+ for BN in (32, 64, 128, 256):
if BN > n: continue
bt = -(-m // BM) * -(-n // BN)
for BK in (256, 512):
⋯ 148 unchanged lines
if best is None:
_L(f" → fallback")
- return {"fallback": True}
+ return {"hot": None}
_L(f" → best={best} @ {best_t:.2f}us")
- return {"fallback": False, "hot": best_go}
+ return {"hot": best_go}
def custom_kernel(data):
- A, B, B_q, B_shuffle, B_scale_sh = data
- m, k = A.shape; n = B.shape[0]
- key = (m, n, k)
+ A = data[0]; Bq = data[2]
+ key = (A.shape[0], Bq.shape[0], A.shape[1])
S = _STATE.get(key)
if S is None:
S = _build(data); _STATE[key] = S
- if S["fallback"]:
- return _ref(A, B_shuffle, B_scale_sh)
- return S["hot"](A, B_q.view(torch.uint8),
- B_scale_sh.view(torch.uint8).reshape(-1))
+ hot = S["hot"]
+ if hot is None:
+ return _ref(A, data[3], data[4])
+ return hot(A, Bq, data[4])
scrolls · 38 diff lines total

Best evidence level for this revision: reported

JSON