Skip to content
KernelIndex
Search⌘K

submission 670896

zaiji100 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:9d3d966a43d739097367380fd9287832ebe9e32fb787d9f6d3821bfc5fd10935
license declaredunknown
license concludedunknown
authorszaiji100
imported2026-08-26

Techniques

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

stages = 1num_warps=meta['nw'], waves_per_eu=0, num_stages=1)

Kernel source

submission.py228 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

from task import input_t, output_t
import os
from typing import Dict, List, Optional, Set, Tuple

os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
os.environ.setdefault("CXX", "clang++")
os.environ.setdefault("AITER_KSPLIT", "1")

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

_fused_mod = None

# ---- Triton fused quant+shuffle (proven fallback) ----
_triton_ok = False
try:
    import triton
    import triton.language as tl
    from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
    _triton_ok = True
except Exception:
    pass

_gemm_asm = getattr(aiter, "gemm_a4w4_asm", None)
_gemm_a4w4 = aiter.gemm_a4w4
_get_padded_m = getattr(aiter, "get_padded_m", None)

if _triton_ok:
    @triton.heuristics({
        "EVEN_M_N": lambda args: args["M"] % args["BLOCK_SIZE_M"] == 0
        and args["N"] % (args["BLOCK_SIZE_N"] * args["NUM_ITER"]) == 0,
    })
    @triton.jit
    def _fused_quant_shuffle_kernel(
        x_ptr, x_fp4_ptr, bs_ptr,
        stride_x_m_in, stride_x_n_in, stride_x_fp4_m_in, stride_x_fp4_n_in,
        M, N,
        SCALE_N: tl.constexpr, SCALE_M_PAD: tl.constexpr, SCALE_N_PAD: tl.constexpr,
        BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
        NUM_ITER: tl.constexpr, NUM_STAGES: tl.constexpr,
        MXFP4_QUANT_BLOCK_SIZE: tl.constexpr, EVEN_M_N: tl.constexpr,
    ):
        pid_m = tl.program_id(0)
        start_n = tl.program_id(1) * NUM_ITER
        stride_x_m = tl.cast(stride_x_m_in, tl.int64)
        stride_x_n = tl.cast(stride_x_n_in, tl.int64)
        stride_x_fp4_m = tl.cast(stride_x_fp4_m_in, tl.int64)
        stride_x_fp4_n = tl.cast(stride_x_fp4_n_in, tl.int64)
        NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
        for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):
            x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
            x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
            x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
            if EVEN_M_N:
                x = tl.load(x_ptr + x_offs, cache_modifier=".cg").to(tl.float32)
            else:
                x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
                x = tl.load(x_ptr + x_offs, mask=x_mask, cache_modifier=".cg").to(tl.float32)
            out_tensor, bs_e8m0 = _mxfp4_quant_op(x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE)
            out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
            out_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
            out_offs = out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
            if EVEN_M_N:
                tl.store(x_fp4_ptr + out_offs, out_tensor)
            else:
                tl.store(x_fp4_ptr + out_offs, out_tensor, mask=(out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :])
            bs_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
            bs_offs_n = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
            num_bs_cols = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
            d0 = bs_offs_m[:, None] // 32
            d1_full = bs_offs_m[:, None] % 32
            d2 = d1_full % 16
            d1 = d1_full // 16
            d3 = bs_offs_n[None, :] // 8
            d4_full = bs_offs_n[None, :] % 8
            d5 = d4_full % 4
            d4 = d4_full // 4
            bs_shuffled_offs = d1 + d4 * 2 + d2 * 4 + d5 * 64 + d3 * 256 + d0 * 32 * SCALE_N_PAD
            bs_valid = (bs_offs_m < M)[:, None] & (bs_offs_n < num_bs_cols)[None, :]
            bs_e8m0 = tl.where(bs_valid, bs_e8m0, 127)
            bs_pad = (bs_offs_m < SCALE_M_PAD)[:, None] & (bs_offs_n < SCALE_N_PAD)[None, :]
            tl.store(bs_ptr + bs_shuffled_offs, bs_e8m0, mask=bs_pad)

def _mk(tile_m, tile_n=128):
    s = f"f4gemm_bf16_per1x32Fp4_BpreShuffle_{tile_m}x{tile_n}"
    return f"_ZN5aiter{len(s)}{s}E"

_K32 = _mk(32)
_ASM_CANDIDATES: Dict[Tuple[int, int, int], List[Tuple[str, int]]] = {
    (4,   2880,  512): [(_K32, 0)],
    (8,   2112, 7168): [(_K32, 0), (_K32, 2)],
    (16,  2112, 7168): [(_K32, 0), (_K32, 2)],
    (16,  3072, 1536): [(_K32, 0)],
    (32,  2880,  512): [(_K32, 0)],
    (32,  4096,  512): [(_K32, 0)],
}

_quant_cache: Dict[Tuple[int, int, int], Tuple[torch.Tensor, torch.Tensor, dict]] = {}
_good_cfg: Dict[Tuple[int, int, int], Tuple[str, int]] = {}
_bad_cfg: Set[Tuple[int, int, int, str, int]] = set()
_padded_rows: Dict[Tuple[int, int, int], int] = {}
_out_bufs: Dict[Tuple[int, int, int], torch.Tensor] = {}
_fused_ok: Optional[bool] = None
_b_raw_cache: Dict[int, torch.Tensor] = {}
_bsc_raw_cache: Dict[int, torch.Tensor] = {}


def _get_raw_bsc(b_sc, N, K):
    ptr = b_sc.data_ptr()
    c = _bsc_raw_cache.get(ptr)
    if c is not None: return c
    s = b_sc.view(torch.uint8)
    sm, sn = s.shape
    s = s.view(sm//32, sn//8, 4, 16, 2, 2).permute(0,5,3,1,4,2).contiguous().view(sm, sn)
    raw = s[:N, :(K+31)//32].contiguous()
    _bsc_raw_cache[ptr] = raw
    return raw


def _ensure_quant(dev_idx, M, K, device):
    key = (dev_idx, M, K)
    cached = _quant_cache.get(key)
    if cached is not None: return cached
    sn_valid = (K + 31) // 32
    sn_pad = ((sn_valid + 7) // 8) * 8
    sm_pad = ((M + 255) // 256) * 256
    if M <= 32:
        bm, bn, ni, nw, ns = triton.next_power_of_2(M), 32, 1, 1, 1
    else:
        ni, bm, bn, nw, ns = 4, 64, 64, 4, 2
        if K <= 16384: bm, bn = 32, 128
    if K <= 1024:
        ni, ns, nw = 1, 1, 4
        bn = max(32, min(256, triton.next_power_of_2(K)))
        bm = min(8, triton.next_power_of_2(M))
    if M <= 32 and K > 1024:
        bm, bn, ni, nw, ns = triton.next_power_of_2(M), 128, 1, 4, 1
    grid = (triton.cdiv(M, bm), triton.cdiv(K, bn * ni))
    x_fp4 = torch.empty((M, K // 2), dtype=torch.uint8, device=device)
    bs = torch.empty((sm_pad, sn_pad), dtype=torch.uint8, device=device)
    meta = {'sn_valid': sn_valid, 'sn_pad': sn_pad, 'sm_pad': sm_pad,
            'bm': bm, 'bn': bn, 'ni': ni, 'nw': nw, 'ns': ns, 'grid': grid}
    _quant_cache[key] = (x_fp4, bs, meta)
    return (x_fp4, bs, meta)

def _get_padded(m, n, k):
    key = (m, n, k)
    r = _padded_rows.get(key)
    if r is not None: return r
    r = m
    if _get_padded_m:
        try: r = int(_get_padded_m(m, n, k, 32))
        except: pass
    _padded_rows[key] = r
    return r


@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
    a = data[0]
    b_sh = data[3]
    b_sc = data[4]
    if not a.is_contiguous():
        a = a.contiguous()

    M, K = a.shape
    N = b_sh.shape[0]
    key = (M, N, K)

    # ---- Triton fused quant+shuffle + ASM GEMM ----
    global _fused_ok
    if _triton_ok and _fused_ok is not False:
        try:
            x_fp4, bs, meta = _ensure_quant(a.device.index or 0, M, K, a.device)
            _fused_quant_shuffle_kernel[meta['grid']](
                a, x_fp4, bs, K, 1, K >> 1, 1,
                M=M, N=K,
                SCALE_N=meta['sn_valid'], SCALE_M_PAD=meta['sm_pad'], SCALE_N_PAD=meta['sn_pad'],
                BLOCK_SIZE_M=meta['bm'], BLOCK_SIZE_N=meta['bn'],
                NUM_ITER=meta['ni'], NUM_STAGES=meta['ns'],
                MXFP4_QUANT_BLOCK_SIZE=32,
                num_warps=meta['nw'], waves_per_eu=0, num_stages=1)
            aq = x_fp4.view(dtypes.fp4x2)
            a_sc = bs.view(dtypes.fp8_e8m0)
            _fused_ok = True
        except Exception:
            _fused_ok = False
            xq, bse = dynamic_mxfp4_quant(a)
            aq = xq.view(dtypes.fp4x2)
            a_sc = e8m0_shuffle(bse).view(dtypes.fp8_e8m0)
    else:
        xq, bse = dynamic_mxfp4_quant(a)
        aq = xq.view(dtypes.fp4x2)
        a_sc = e8m0_shuffle(bse).view(dtypes.fp8_e8m0)

    m, n, k = M, N, K
    if _gemm_asm is not None:
        cfg = _good_cfg.get(key)
        if cfg is None:
            for kn, sk in _ASM_CANDIDATES.get(key, []):
                if (m, n, k, kn, sk) not in _bad_cfg:
                    cfg = (kn, sk)
                    break
        if cfg is not None:
            kn, sk = cfg
            pr = _get_padded(m, n, k)
            buf_key = (aq.device.index or 0, pr, n)
            out = _out_bufs.get(buf_key)
            if out is None:
                out = torch.empty((pr, n), dtype=torch.bfloat16, device=aq.device)
                _out_bufs[buf_key] = out
            try:
                if sk > 0: out.zero_()
                _gemm_asm(aq.view(m, -1), b_sh, a_sc.view(m, -1), b_sc,
                           out, kn, None, 1.0, 0.0, True, sk)
                _good_cfg[key] = cfg
                return out[:m]
            except Exception:
                _bad_cfg.add((m, n, k, kn, sk))

    return _gemm_a4w4(aq, b_sh, a_sc, b_sc, dtype=dtypes.bf16, bpreshuffle=True)
scrolls · 228 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