Skip to content
KernelIndex
Search⌘K

submission 753771

Aniket Sadashiva · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:0d80c6ce6866c25a1ddda3824b0a9ac0b112305486157e334dc1300347834d20
license declaredunknown
license concludedunknown
authorsAniket Sadashiva
imported2026-08-15

Techniques

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

num-warps = 4vh366: vh362 base + quant kernel tuning (num_warps=4 instead of 2).
tile-k = 256BM=16, BN=64, BK=256, GM=1, NKS=1, SBS=kp<<1, EK=True,
tile-m = 16BM=16, BN=64, BK=256, GM=1, NKS=1, SBS=kp<<1, EK=True,
tile-n = 64BM=16, BN=64, BK=256, GM=1, NKS=1, SBS=kp<<1, EK=True,

Kernel source

submission_vh366.py237 lines
"""
vh366: vh362 base + quant kernel tuning (num_warps=4 instead of 2).
Also try BS=32 for the quant kernel (more blocks for M=64).
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t

import aiter
from aiter import dtypes
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op

try:
    import triton._utils as _tu
    _d = _tu.type_canonicalisation_dict
    _d.setdefault("float4_e2m1fn_x2", "u8")
    _d.setdefault("float8_e8m0fnu", "u8")
    _d.setdefault("float4_e2m1fn", "u8")
except Exception:
    pass

try:
    from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
    HAS_ASM = True
except ImportError:
    gemm_a4w4_asm = None
    HAS_ASM = False

_KN = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"

@triton.jit
def _pg(pid: int, npm: int, npn: int, GM: tl.constexpr = 1):
    if GM == 1: return pid // npn, pid % npn
    nig = GM * npn; gid = pid // nig; fpm = gid * GM
    gsm = min(npm - fpm, GM); tl.assume(gsm >= 0)
    return fpm + (pid % gsm), (pid % nig) // gsm

@triton.jit
def _gk(ap, bp, cp, bsp, M, N, K,
    sam, sak, sbk, sbn, sck, scm, scn, sbsk, sbsn,
    BM: tl.constexpr, BN: tl.constexpr, BK: tl.constexpr, GM: tl.constexpr,
    NKS: tl.constexpr, SBS: tl.constexpr, EK: tl.constexpr,
    PRESHUFFLE: tl.constexpr):
    tl.assume(sam > 0); tl.assume(sak > 0); tl.assume(sbk > 0); tl.assume(sbn > 0)
    tl.assume(scm > 0); tl.assume(scn > 0); tl.assume(sbsk > 0); tl.assume(sbsn > 0)
    pu = tl.program_id(0); pk = pu % NKS; p = pu // NKS
    npm = tl.cdiv(M, BM); npn = tl.cdiv(N, BN)
    if NKS == 1: pm, pn = _pg(p, npm, npn, GM=GM)
    else: pm = p // npn; pn = p % npn
    tl.assume(pm >= 0); tl.assume(pn >= 0); tl.assume(pk >= 0)
    SG: tl.constexpr = 32; ST: tl.constexpr = BK // SG
    if (pk * SBS // 2) < K:
        nki = tl.cdiv(SBS // 2, BK // 2)
        okb = tl.arange(0, BK); oksb = pk * SBS + okb
        oam = (pm * BM + tl.arange(0, BM)) % M
        apt = ap + (oam[:, None] * sam + oksb[None, :] * sak)
        if PRESHUFFLE:
            obn_ps = (pn * (BN // 16) + tl.arange(0, BN // 16)) % (N // 16)
            oks_ps = pk * (SBS // 2) * 16 + tl.arange(0, (BK // 2) * 16)
            bpt = bp + obn_ps[:, None] * sbn + oks_ps[None, :] * sbk
            obsn = (pn * (BN // 32) + tl.arange(0, BN // 32)) % (N // 32)
            obsk = (pk * (SBS // SG) * 32) + tl.arange(0, BK // SG * 32)
            bspt = bsp + obsn[:, None] * sbsn + obsk[None, :] * sbsk
        else:
            ok = tl.arange(0, BK // 2); oks_nat = pk * (SBS // 2) + ok
            obn_nat = (pn * BN + tl.arange(0, BN)) % N
            bpt = bp + (obn_nat[:, None] * sbn + oks_nat[None, :] * sbk)
            ok2 = pk * (SBS // SG) + tl.arange(0, ST)
            d0 = obn_nat // 32; d1 = (obn_nat & 31) >> 4; d2 = obn_nat & 15
            srp = d0 * (32 * sbsn) + d2 * 4 + d1
            d3 = ok2 >> 3; d4 = (ok2 & 7) >> 2; d5 = ok2 & 3
            scp2 = d3 * 256 + d5 * 64 + d4 * 2
            bso = srp[:, None] + scp2[None, :]
        acc = tl.zeros((BM, BN), dtype=tl.float32)
        for ki in range(0, nki):
            if PRESHUFFLE:
                bs = (tl.load(bspt)
                    .reshape(BN // 32, BK // SG // 8, 4, 16, 2, 2, 1)
                    .permute(0, 5, 3, 1, 4, 2, 6)
                    .reshape(BN, BK // SG))
                if EK: ab = tl.load(apt); b_raw = tl.load(bpt)
                else:
                    ab = tl.load(apt, mask=okb[None, :] < SBS, other=0)
                    b_raw = tl.load(bpt, mask=(obn_ps[:, None] < (N // 16)) & (oks_ps[None, :] < (K * 16)), other=0)
                b = (b_raw.reshape(1, BN // 16, BK // 64, 2, 16, 16)
                    .permute(0, 1, 4, 2, 3, 5)
                    .reshape(BN, BK // 2)
                    .trans(1, 0))
            else:
                aki = pk * nki + ki
                bs = tl.load(bsp + bso)
                if EK: ab = tl.load(apt); b = tl.load(bpt).trans(1, 0)
                else:
                    ab = tl.load(apt, mask=okb[None, :] < 2 * K - aki * BK, other=0)
                    b = tl.load(bpt, mask=tl.arange(0, BK // 2)[None, :] < K - aki * (BK // 2), other=0).trans(1, 0)
            a, asc = _mxfp4_quant_op(ab, BK, BM, 32)
            acc += tl.dot_scaled(a, asc, "e2m1", b, bs, "e2m1")
            apt += BK * sak
            if PRESHUFFLE:
                bpt += (BK // 2) * 16 * sbk; bspt += BK * sbsk
            else:
                bpt += (BK // 2) * sbk; ok2 += ST
                d3 = ok2 >> 3; d4 = (ok2 & 7) >> 2; d5 = ok2 & 3
                scp2 = d3 * 256 + d5 * 64 + d4 * 2
                bso = srp[:, None] + scp2[None, :]
        c = acc.to(cp.type.element_ty)
        ocm = pm * BM + tl.arange(0, BM).to(tl.int64)
        ocn = pn * BN + tl.arange(0, BN).to(tl.int64)
        cpt = cp + scm * ocm[:, None] + scn * ocn[None, :] + pk * sck
        cm = (ocm[:, None] < M) & (ocn[None, :] < N)
        if NKS == 1: tl.store(cpt, c, mask=cm, cache_modifier=".wt")
        else: tl.store(cpt, c, mask=cm)

@triton.jit
def _rk(src, dst, M, N, ss, sm, sn, dm, dn,
        NKS: tl.constexpr, BM: tl.constexpr, BN: tl.constexpr):
    pid = tl.program_id(0); npn = tl.cdiv(N, BN); pm = pid // npn; pn = pid % npn
    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 s in tl.static_range(NKS):
        a += tl.load(src + s * ss + om[:, None] * sm + on[None, :] * sn, mask=mk, other=0.0)
    tl.store(dst + om[:, None] * dm + on[None, :] * dn, a.to(tl.bfloat16), mask=mk, cache_modifier=".wt")

@triton.jit
def _qk(xp, fp, bp, sxm, sxn, sfm, sfn, M, N, scN, sMp, sNp, BS: tl.constexpr):
    qb: tl.constexpr = 32; pm = tl.program_id(0); pn = tl.program_id(1)
    sxm64 = tl.cast(sxm, tl.int64); sxn64 = tl.cast(sxn, tl.int64)
    sfm64 = tl.cast(sfm, tl.int64); sfn64 = tl.cast(sfn, tl.int64)
    xom = pm * BS + tl.arange(0, BS); xon = pn * qb + tl.arange(0, qb)
    xmk = (xom < M)[:, None] & (xon < N)[None, :]
    x = tl.load(xp + xom[:, None] * sxm64 + xon[None, :] * sxn64, mask=xmk).to(tl.float32)
    xf, be = _mxfp4_quant_op(x, qb, BS, qb)
    oon = pn * (qb // 2) + tl.arange(0, qb // 2)
    omk = (xom < M)[:, None] & (oon < (N // 2))[None, :]
    tl.store(fp + xom[:, None] * sfm64 + oon[None, :] * sfn64, xf, mask=omk)
    bv = tl.reshape(be, [BS]); bm = xom; bn = pn
    d0 = bm // 32; r32 = bm % 32; d2 = r32 % 16; d1 = r32 // 16
    d3 = bn // 8; r8 = bn % 8; d5 = r8 % 4; d4 = r8 // 4
    so = d1 + d4 * 2 + d2 * 4 + d5 * 64 + d3 * 256 + d0 * (32 * scN)
    m1 = (bm < M) & (bn < scN); m2 = (bm < sMp) & (bn < sNp)
    bv = tl.where(m1, bv, 127); tl.store(bp + so, bv, mask=m2)


_C = {}; _W = set()
def _dk(d): return d.type, d.index
def _ct(n, s, dt, d):
    k = (n, _dk(d)); o = _C.get(k)
    if o is None: o = torch.empty(s, dtype=dt, device=d); _C[k] = o
    return o
def _ab(m, n, k, d):
    key = (("a", m, n, k), _dk(d)); c = _C.get(key)
    if c: return c
    sv = triton.cdiv(k, 32); sp = triton.cdiv(sv, 8) * 8; sm = triton.cdiv(m, 32) * 32
    xf = torch.empty((m, k // 2), dtype=torch.uint8, device=d)
    bs = torch.empty((sm, sp), dtype=torch.uint8, device=d)
    out = torch.empty((((m + 31) // 32) * 32, n), dtype=torch.bfloat16, device=d)
    xf_fp4 = xf.view(dtypes.fp4x2).view(m, k // 2); bs_e8m0 = bs.view(dtypes.fp8_e8m0)
    c = (xf, bs, xf_fp4, bs_e8m0, out, sv, sm, sp); _C[key] = c; return c


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]
    kp = k >> 1

    if k == 512:
        b_q_u8 = b_q.view(torch.uint8)
        b_scale_u8 = b_scale_sh.view(torch.uint8)
        out = _ct(("o", m, n, k), (m, n), torch.bfloat16, a.device)
        g = triton.cdiv(m, 16) * triton.cdiv(n, 64)
        _gk[(g,)](a, b_q_u8, out, b_scale_u8, m, n, kp,
            a.stride(0), a.stride(1), b_q_u8.stride(1), b_q_u8.stride(0),
            0, out.stride(0), out.stride(1), 1, b_scale_u8.shape[1],
            BM=16, BN=64, BK=256, GM=1, NKS=1, SBS=kp<<1, EK=True,
            PRESHUFFLE=False, num_warps=4)
        return out

    if m == 16 and k == 7168:
        bsu8 = b_shuffle.view(torch.uint8)
        bt = bsu8.view(n // 16, kp * 16)
        bscu8 = b_scale_sh.view(torch.uint8)
        bst = bscu8.reshape(bscu8.shape[0] // 32, bscu8.shape[1] * 32)
        NS = 7; SBS = (2 * kp) // NS
        skb = _ct(("sk", m, n, k), (NS, m, n), torch.float32, a.device)
        out = _ct(("so", m, n, k), (m, n), torch.bfloat16, a.device)
        g = triton.cdiv(n, 128)
        _gk[(g * NS,)](a, bt, skb, bst, m, n, kp,
            a.stride(0), a.stride(1), bt.stride(1), bt.stride(0),
            m * n, skb.stride(1), skb.stride(2), bst.stride(1), bst.stride(0),
            BM=16, BN=128, BK=512, GM=1, NKS=NS, SBS=SBS, EK=True,
            PRESHUFFLE=True, num_warps=4)
        _rk[(g,)](skb, out, m, n, skb.stride(0), skb.stride(1), skb.stride(2),
            out.stride(0), out.stride(1), NKS=NS, BM=16, BN=128, num_warps=4)
        return out

    # M=64/M=256: quant with num_warps=4 (was 2) + ASM
    if HAS_ASM:
        xf, bs, xf_fp4, bs_e8m0, out, sv, sm, sp = _ab(m, n, k, a.device)
        _qk[(((m + 63) >> 6), sv)](a, xf, bs,
            a.stride(0), a.stride(1), xf.stride(0), xf.stride(1),
            M=m, N=k, scN=sv, sMp=sm, sNp=sp, BS=64, num_warps=4)
        gemm_a4w4_asm(xf_fp4, b_shuffle, bs_e8m0, b_scale_sh, out, _KN,
            bpreshuffle=True, log2_k_split=1)
        return out[:m]

    from aiter.ops.triton.quant import dynamic_mxfp4_quant
    from aiter.utility.fp4_utils import e8m0_shuffle
    afp4, asc = dynamic_mxfp4_quant(a)
    return aiter.gemm_a4w4(afp4.view(dtypes.fp4x2), b_shuffle,
        e8m0_shuffle(asc).view(dtypes.fp8_e8m0), b_scale_sh, dtype=dtypes.bf16, bpreshuffle=True)


def _ms(m, n, k, dev):
    from aiter.ops.shuffle import shuffle_weight
    a = torch.zeros((m, k), dtype=torch.bfloat16, device=dev)
    b = torch.zeros((n, k), dtype=torch.bfloat16, device=dev)
    bq = torch.zeros((n, k // 2), dtype=torch.uint8, device=dev).view(dtypes.fp4x2)
    sv = triton.cdiv(k, 32); sp = triton.cdiv(sv, 8) * 8; sm = triton.cdiv(n, 32) * 32
    bss = torch.full((sm, sp), 127, dtype=torch.uint8, device=dev).view(dtypes.fp8_e8m0)
    return a, b, bq, shuffle_weight(bq, layout=(16, 16)), bss

def _pw():
    if not torch.cuda.is_available(): return
    dev = torch.device("cuda"); dk = _dk(dev)
    if dk in _W: return
    _W.add(dk)
    for m, n, k in [(4,2880,512),(32,4096,512),(32,2880,512),(16,2112,7168),(64,7168,2048),(256,3072,1536)]:
        try: custom_kernel(_ms(m, n, k, dev))
        except Exception as e: print(f"[vh366] prewarm ({m},{n},{k}): {e}", flush=True)
    try: torch.cuda.synchronize(dev)
    except: pass

try: _pw()
except: pass
scrolls · 237 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