Skip to content
KernelIndex
Search⌘K

submission 627274

Bortlesboat · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

v78_splitk0.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-627274?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
8.94µs
#99 of 1143
2026-03-25

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f50d65705980d918670afdd84686a443d0865fca357cd03cebfc093a68b30301
license declaredunknown
license concludedunknown
authorsBortlesboat
imported2026-08-15

Techniques

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

autotune@triton.autotune(
num-warps = 4triton.Config({'BSM': 4, 'BSN': 512, 'NI': 1, 'NS': 1}, num_warps=4),
split-k0, out.stride(0), out.stride(1), # stride_ck=0 for no split-K

Kernel source

v78_splitk0.py342 lines
"""v53: Use aiter's _gemm_a16wfp4_preshuffle_kernel directly with PREQUANT=True.

Based on dgavriloff/amd-structkernel approach (8.667μs proven).
Key insight: the preshuffle kernel handles B_shuffle layout natively with
coalesced loads + compile-time reshape/permute/trans. No manual B transposition needed.

Shapes 1-5 (M<=64): _gemm_a16wfp4_preshuffle_kernel with PREQUANT=True
Shape 6 (M=256): quant A + ASM gemm_a4w4_asm (32x128 tile)
"""
from task import input_t, output_t
import os, sys
os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
os.environ["HSA_TOOLS_LIB"] = ""

import torch
import triton
import triton.language as tl
from aiter import dtypes
import aiter
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle

_fp4x2 = dtypes.fp4x2; _fp8_e8m0 = dtypes.fp8_e8m0; _bf16 = dtypes.bf16

# Try importing aiter's preshuffle kernel
_HAS_PRESHUFFLE = False
try:
    from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (
        _gemm_a16wfp4_preshuffle_kernel,
    )
    _HAS_PRESHUFFLE = True
except ImportError:
    print("[v53] WARN: _gemm_a16wfp4_preshuffle_kernel not available", file=sys.stderr)

# Try importing gluon reduce kernel for split-K
_HAS_GLUON_REDUCE = False
try:
    from aiter.ops.triton.gluon.gemm_afp4wfp4 import (
        _gemm_afp4wfp4_reduce_kernel as _gluon_reduce_kernel,
    )
    _HAS_GLUON_REDUCE = True
except ImportError:
    pass

if not _HAS_GLUON_REDUCE:
    try:
        from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
            _gemm_afp4wfp4_reduce_kernel as _gluon_reduce_kernel,
        )
        _HAS_GLUON_REDUCE = True
    except ImportError:
        print("[v53] WARN: reduce kernel not available", file=sys.stderr)

# Try importing _mxfp4_quant_op for fallback fused quant
try:
    from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
    _HAS_IQ = True
except ImportError:
    _HAS_IQ = False

# ---- Shape-specific configs (from dgavriloff's proven v224) ----
# (M, N, K) -> (BSM, BSN, BSK, warps, stages, waves_per_eu, cache, split_k)
SHAPE_CONFIGS = {
    (4, 2880, 512):   (4, 128, 256, 4, 2, 0, ".cg", 1),
    (16, 2112, 7168): (8, 128, 256, 4, 2, 2, ".cg", 7),
    (32, 4096, 512):  (8, 128, 256, 4, 2, 2, None, 1),
    (32, 2880, 512):  (8, 128, 256, 4, 2, 2, None, 1),
    (64, 7168, 2048): (16, 128, 256, 4, 2, 2, ".cg", 1),
}

# ---- Buffer caches ----
_ob = {}; _ws = {}; _bw = {}; _bsc = {}; _fo = {}

def _get_obuf(M, N, d):
    k = (M, N)
    if k not in _ob:
        _ob[k] = torch.empty((M, N), dtype=torch.bfloat16, device=d)
    return _ob[k]

def _reshape_b(B_shuffle, N, K):
    """Reshape B_shuffle for preshuffle kernel: (N//16, (K//2)*16)"""
    k = (N, K)
    if k in _bw:
        bw, ref = _bw[k]
        if ref.data_ptr() == B_shuffle.data_ptr():
            return bw
    bw = B_shuffle.view(torch.uint8).reshape(N // 16, (K // 2) * 16)
    _bw[k] = (bw, B_shuffle)
    return bw

def _reshape_bsc(B_scale_sh, N, K):
    """Reshape B_scale_sh for preshuffle kernel: (bs0//32, bs1*32)"""
    k = (N, K)
    if k in _bsc:
        bsc, ref = _bsc[k]
        if ref.data_ptr() == B_scale_sh.data_ptr():
            return bsc
    bs = B_scale_sh.view(torch.uint8)
    bs0, bs1 = bs.shape
    bsc = bs.reshape(bs0 // 32, bs1 * 32)
    _bsc[k] = (bsc, B_scale_sh)
    return bsc

def _run_preshuffle(A, B_shuffle, B_scale_sh, M, N, K, cfg):
    BSM, BSN, BSK, warps, stages, wpe, cache, num_ksplit = cfg
    B_w = _reshape_b(B_shuffle, N, K)
    B_sc = _reshape_bsc(B_scale_sh, N, K)
    K_kernel = K // 2

    if num_ksplit == 1:
        out = _get_obuf(M, N, A.device)
        grid_size = triton.cdiv(M, BSM) * triton.cdiv(N, BSN)
        _gemm_a16wfp4_preshuffle_kernel[(grid_size,)](
            A, B_w, out, B_sc,
            M, N, K_kernel,
            A.stride(0), A.stride(1),
            B_w.stride(0), B_w.stride(1),
            0, out.stride(0), out.stride(1),  # stride_ck=0 for no split-K
            B_sc.stride(0), B_sc.stride(1),
            BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=BSN, BLOCK_SIZE_K=BSK,
            GROUP_SIZE_M=1,
            NUM_KSPLIT=1,
            SPLITK_BLOCK_SIZE=2 * K_kernel,
            num_warps=warps, num_stages=stages, waves_per_eu=wpe,
            matrix_instr_nonkdim=16,
            PREQUANT=True,
            cache_modifier=cache,
        )
        return out
    else:
        # Split-K path
        wk = (num_ksplit, M, N)
        if wk not in _ws:
            _ws[wk] = torch.empty(wk, dtype=torch.float32, device=A.device)
        y_pp = _ws[wk]

        # Compute SPLITK_BLOCK_SIZE: each split handles K_kernel/num_ksplit elements
        # SPLITK_BLOCK_SIZE = ceil(K_kernel / num_ksplit) rounded up to BSK
        k_per_split = (K_kernel + num_ksplit - 1) // num_ksplit
        SPLITK_BLOCK_SIZE = ((k_per_split + BSK - 1) // BSK) * BSK * 2

        grid_size = num_ksplit * triton.cdiv(M, BSM) * triton.cdiv(N, BSN)
        _gemm_a16wfp4_preshuffle_kernel[(grid_size,)](
            A, B_w, y_pp, B_sc,
            M, N, K_kernel,
            A.stride(0), A.stride(1),
            B_w.stride(0), B_w.stride(1),
            y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
            B_sc.stride(0), B_sc.stride(1),
            BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=BSN, BLOCK_SIZE_K=BSK,
            GROUP_SIZE_M=1,
            NUM_KSPLIT=num_ksplit,
            SPLITK_BLOCK_SIZE=SPLITK_BLOCK_SIZE,
            num_warps=warps, num_stages=stages, waves_per_eu=wpe,
            matrix_instr_nonkdim=16,
            PREQUANT=True,
            cache_modifier=cache,
        )

        # Reduce split-K partials
        out = _get_obuf(M, N, A.device)
        ACTUAL_KSPLIT = triton.cdiv(K_kernel, SPLITK_BLOCK_SIZE // 2)

        if _HAS_GLUON_REDUCE:
            reduce_grid = (triton.cdiv(M, 16), triton.cdiv(N, 64))
            _gluon_reduce_kernel[reduce_grid](
                y_pp, out,
                M, N,
                y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
                out.stride(0), out.stride(1),
                16, 64,
                ACTUAL_KSPLIT,
                triton.next_power_of_2(num_ksplit),
            )
        else:
            # Fallback: simple sum in Python
            out.copy_(y_pp.sum(dim=0).to(torch.bfloat16))
        return out

# ---- A quant for ASM path (M=256) ----
@triton.autotune(
    configs=[
        triton.Config({'BSM': 4, 'BSN': 512, 'NI': 1, 'NS': 1}, num_warps=4),
        triton.Config({'BSM': 8, 'BSN': 256, 'NI': 1, 'NS': 1}, num_warps=4),
        triton.Config({'BSM': 8, 'BSN': 512, 'NI': 1, 'NS': 1}, num_warps=4),
        triton.Config({'BSM': 16, 'BSN': 128, 'NI': 4, 'NS': 2}, num_warps=4),
        triton.Config({'BSM': 32, 'BSN': 512, 'NI': 1, 'NS': 1}, num_warps=4),
        triton.Config({'BSM': 64, 'BSN': 128, 'NI': 4, 'NS': 2}, num_warps=4),
    ],
    key=['M', 'N'],
)
@triton.jit
def _quant_kernel(x, fp, bs, sxm, sxn, sfm, sfn, M, N, snp,
                  BSM: tl.constexpr, BSN: tl.constexpr,
                  NI: tl.constexpr, NS: tl.constexpr, QBS: tl.constexpr):
    pm = tl.program_id(0); sn = tl.program_id(1) * NI
    xm = tl.cast(sxm, tl.int64); xn = tl.cast(sxn, tl.int64)
    fm = tl.cast(sfm, tl.int64); fn = tl.cast(sfn, tl.int64)
    NQB: tl.constexpr = BSN // QBS
    for pn in tl.range(sn, min(sn + NI, N), num_stages=NS):
        om = pm * BSM + tl.arange(0, BSM)
        on = pn * BSN + tl.arange(0, BSN)
        v = tl.load(x + om[:, None] * xm + on[None, :] * xn,
                    mask=(om < M)[:, None] & (on < N)[None, :], other=0.0,
                    cache_modifier=".cg").to(tl.float32)
        f4, sc = _mxfp4_quant_op(v, BSN, BSM, QBS)
        fo = pm * BSM + tl.arange(0, BSM)
        fno = pn * BSN // 2 + tl.arange(0, BSN // 2)
        tl.store(fp + fo[:, None] * fm + fno[None, :] * fn, f4,
                 mask=(fo < M)[:, None] & (fno < N // 2)[None, :])
        sm = pm * BSM + tl.arange(0, BSM)
        sk = pn * NQB + tl.arange(0, NQB)
        d0 = sm // 32; d1 = (sm % 32) // 16; d2 = sm % 16
        d3 = sk // 8; d4 = (sk % 8) // 4; d5 = sk % 4
        fl = ((d0 * (snp * 32))[:, None] + (d3 * 256)[None, :]
              + (d5 * 64)[None, :] + (d2 * 4)[:, None]
              + (d4 * 2)[None, :] + d1[:, None])
        tl.store(bs + fl, sc,
                 mask=(sm < M)[:, None] & (sk < N // QBS)[None, :])

_qb = {}

def _get_qbuf(M, K, d):
    k = (M, K)
    if k not in _qb:
        nq = K // 32; sp = ((M + 255) // 256) * 256; snp = ((nq + 7) // 8) * 8
        _qb[k] = (torch.empty((M, K // 2), dtype=torch.uint8, device=d),
                  torch.zeros(sp * snp, dtype=torch.uint8, device=d), sp, snp)
    return _qb[k]

def _get_snp(K):
    return ((K // 32 + 7) // 8) * 8

def _quant_a(A):
    M, K = A.shape
    f, b, sp, snp = _get_qbuf(M, K, A.device)
    g = lambda m: (triton.cdiv(M, m['BSM']), triton.cdiv(K, m['BSN'] * m['NI']))
    _quant_kernel[g](A, f, b, A.stride(0), A.stride(1),
                     f.stride(0), f.stride(1), M, K, snp, QBS=32)
    return f.view(_fp4x2), b.view(sp, snp).view(_fp8_e8m0)

def _kn(t):
    i = f"f4gemm_bf16_per1x32Fp4_BpreShuffle_{t}"
    return f"_ZN5aiter{len(i)}{i}E"

_ga = None; _gg = None; _ao = {}

def custom_kernel(data: input_t) -> output_t:
    global _ga, _gg
    A, B, Bq, Bs, Bss = data  # Bs=B_shuffle, Bss=B_scale_sh
    A = A.contiguous()
    M, K = A.shape; N = B.shape[0]
    key = (M, N, K)

    # ---- PRESHUFFLE KERNEL: M<=64 (shapes 1-5) ----
    if _HAS_PRESHUFFLE and M <= 64:
        cfg = SHAPE_CONFIGS.get(key)
        if cfg is None:
            # Default config for unknown shapes
            if K > 2048:
                cfg = (8, 128, 256, 4, 2, 2, ".cg", max(1, K // 1024))
            else:
                cfg = (min(M, 16), 128, 256, 4, 2, 2, ".cg", 1)

        pk = ("ps", key)
        if pk not in _fo:
            try:
                C = _run_preshuffle(A, Bs, Bss, M, N, K, cfg)
                # Verify correctness on first call
                if _gg is None: _gg = aiter.gemm_a4w4
                Aq, Ac = _quant_a(A) if _HAS_IQ else _sq(A)
                ref = _gg(Aq, Bs, Ac, Bss, dtype=_bf16, bpreshuffle=True)
                me = (C.float() - ref.float()).abs().max().item()
                _fo[pk] = me < 2.0
                if _fo[pk]:
                    print(f"[v53] PRESHUFFLE OK {key} err={me:.1f}", file=sys.stderr)
                    return C
                else:
                    print(f"[v53] PRESHUFFLE FAIL {key} err={me:.1f}", file=sys.stderr)
            except Exception as e:
                _fo[pk] = False
                print(f"[v53] PRESHUFFLE ERR {key}: {e}", file=sys.stderr)
        elif _fo[pk]:
            return _run_preshuffle(A, Bs, Bss, M, N, K, cfg)

    # ---- M=256: quant A + ASM GEMM (two-phase) ----
    Aq, Ac = _quant_a(A) if _HAS_IQ else _sq(A)
    if _gg is None: _gg = aiter.gemm_a4w4
    if _ga is None: _ga = aiter.gemm_a4w4_asm

    if M == 256 and K >= 1536:
        tile = "32x128"
        sp = None  # no split-K (aiter get_GEMM_config recommends splitK=0)
        tk = ("asm256", tile)
        if tk not in _ao:
            try:
                out = _get_obuf(M, N, A.device)
                _ga(Aq, Bs, Ac, Bss, out, _kn(tile), bpreshuffle=True, log2_k_split=sp)
                _ao[tk] = True
            except Exception:
                _ao[tk] = False
        if _ao.get(tk, False):
            out = _get_obuf(M, N, A.device)
            return _ga(Aq, Bs, Ac, Bss, out, _kn(tile), bpreshuffle=True, log2_k_split=sp)

    if M == 64:
        tile = "32x128"
        sp = None  # try no split-K
        tk = ("asm64", tile)
        if tk not in _ao:
            try:
                out = _get_obuf(M, N, A.device)
                _ga(Aq, Bs, Ac, Bss, out, _kn(tile), bpreshuffle=True, log2_k_split=sp)
                _ao[tk] = True
            except Exception:
                _ao[tk] = False
        if _ao.get(tk, False):
            out = _get_obuf(M, N, A.device)
            return _ga(Aq, Bs, Ac, Bss, out, _kn(tile), bpreshuffle=True, log2_k_split=sp)

    if M <= 32:
        tile = "32x128"
        sp = 1 if K >= 2048 else None
        if tile not in _ao:
            try:
                out = _get_obuf(M, N, A.device)
                _ga(Aq, Bs, Ac, Bss, out, _kn(tile), bpreshuffle=True, log2_k_split=sp)
                _ao[tile] = True
            except Exception:
                _ao[tile] = False
        if _ao.get(tile, False):
            out = _get_obuf(M, N, A.device)
            return _ga(Aq, Bs, Ac, Bss, out, _kn(tile), bpreshuffle=True, log2_k_split=sp)

    # ---- FALLBACK: gemm_a4w4 ----
    return _gg(Aq, Bs, Ac, Bss, dtype=_bf16, bpreshuffle=True)

def _sq(A):
    Aq, Ac = dynamic_mxfp4_quant(A)
    Ac = e8m0_shuffle(Ac)
    return Aq.view(_fp4x2), Ac.view(_fp8_e8m0)
scrolls · 342 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