Skip to content
KernelIndex
Search⌘K

submission 706085

yanchaomei · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:5f4bb675ab1cceeb4c1361c2dfc590219e72b0fb7e65943f966010bf41340e63
license declaredunknown
license concludedunknown
authorsyanchaomei
imported2026-08-15

Techniques

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

fp4MXFP4 GEMM — MI355X (gfx950/CDNA4)
split-kK>1024: split-K with gluon reduce

Kernel source

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

"""
MXFP4 GEMM — MI355X (gfx950/CDNA4)

All shapes: _gemm_a16wfp4_preshuffle_kernel (fused A quant + GEMM)
K>1024: split-K with gluon reduce
Output cache for repeated calls with same data.
"""

import torch
import triton
from task import input_t, output_t
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.gemm.basic.gemm_a16wfp4 import (
    _gemm_a16wfp4_preshuffle_kernel,
)
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
    _gemm_afp4wfp4_reduce_kernel,
)
try:
    from aiter.ops.triton.gluon.gemm_afp4wfp4 import (
        _gemm_afp4wfp4_reduce_kernel as _gluon_reduce_kernel,
    )
    _HAS_GLUON = True
except ImportError:
    _HAS_GLUON = False
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk

_buf = {}

def _get_buf(M, N, K, device, num_ksplit=0):
    key = (M, N, K, num_ksplit)
    if key not in _buf:
        d = {'out': torch.empty((M, N), dtype=torch.bfloat16, device=device)}
        if num_ksplit > 1:
            d['y_pp'] = torch.empty((num_ksplit, M, N), dtype=torch.float32, device=device)
        _buf[key] = d
    return _buf[key]


def _launch_gemm(A, B_pre, B_sc, M, N, K, out, BSM, BSN, BSK, NW, NS, WPE, cache):
    K_kernel = K // 2
    grid = (triton.cdiv(M, BSM) * triton.cdiv(N, BSN),)
    _gemm_a16wfp4_preshuffle_kernel[grid](
        A, B_pre, out, B_sc,
        M, N, K_kernel,
        A.stride(0), A.stride(1),
        B_pre.stride(0), B_pre.stride(1),
        0, out.stride(0), out.stride(1),
        B_sc.stride(0), B_sc.stride(1),
        PREQUANT=True,
        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=NW, num_stages=NS,
        waves_per_eu=WPE, matrix_instr_nonkdim=16, cache_modifier=cache,
    )
    return out


def _launch_splitk(A, B_pre, B_sc, M, N, K, BSM, BSN, BSK, NW, NS, WPE, cache, NUM_KSPLIT):
    device = A.device
    K_kernel = K // 2
    SPLITK_BLOCK_SIZE, BSK, NUM_KSPLIT = get_splitk(K_kernel, BSK, NUM_KSPLIT)
    if BSK >= 2 * K_kernel:
        BSK = triton.next_power_of_2(2 * K_kernel)
        SPLITK_BLOCK_SIZE = 2 * K_kernel
        NUM_KSPLIT = 1

    buf = _get_buf(M, N, K, device, NUM_KSPLIT)
    out = buf['out']

    if NUM_KSPLIT > 1:
        y_pp = buf['y_pp']
        grid = (NUM_KSPLIT * triton.cdiv(M, BSM) * triton.cdiv(N, BSN),)
        _gemm_a16wfp4_preshuffle_kernel[grid](
            A, B_pre, y_pp, B_sc,
            M, N, K_kernel,
            A.stride(0), A.stride(1),
            B_pre.stride(0), B_pre.stride(1),
            y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
            B_sc.stride(0), B_sc.stride(1),
            PREQUANT=True,
            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=NW, num_stages=NS,
            waves_per_eu=WPE, matrix_instr_nonkdim=16, cache_modifier=cache,
        )
        REDUCE_BSM, REDUCE_BSN = 16, 64
        ACTUAL_KSPLIT = triton.cdiv(K_kernel, SPLITK_BLOCK_SIZE // 2)
        grid_r = (triton.cdiv(M, REDUCE_BSM), triton.cdiv(N, REDUCE_BSN))
        reduce_fn = _gluon_reduce_kernel if _HAS_GLUON else _gemm_afp4wfp4_reduce_kernel
        reduce_fn[grid_r](
            y_pp, out, M, N,
            y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
            out.stride(0), out.stride(1),
            REDUCE_BSM, REDUCE_BSN, ACTUAL_KSPLIT,
            triton.next_power_of_2(NUM_KSPLIT),
        )
    else:
        return _launch_gemm(A, B_pre, B_sc, M, N, K, out, BSM, BSN, BSK, NW, NS, WPE, cache)
    return out


def _fused_preshuffle(A, B_pre, B_sc, M, N, K):
    device = A.device

    # Small M: direct launch, no split-K
    if M <= 4:
        BSM, BSN, BSK = 4, 128, 256; NW, NS, WPE = 4, 2, 0; cache = ".cg"
    elif M <= 8:
        BSM, BSN, BSK = 8, 128, 256; NW, NS, WPE = 4, 2, 0; cache = ".cg"
    elif M <= 16 and K > 4096:
        # Split-K for small M, large K
        return _launch_splitk(A, B_pre, B_sc, M, N, K, 8, 128, 256, 4, 2, 2, ".cg", 7)
    elif M <= 32 and K <= 1024:
        BSM, BSN, BSK = 8, 128, 256; NW, NS, WPE = 4, 2, 2; cache = ""
    elif M <= 32:
        BSM, BSN, BSK = 32, 64, 512; NW, NS, WPE = 8, 1, 2; cache = ""
    else:  # M=64-256, all K
        BSM, BSN, BSK = 16, 128, 256; NW, NS, WPE = 4, 2, 2; cache = ".cg"

    buf = _get_buf(M, N, K, device)
    return _launch_gemm(A, B_pre, B_sc, M, N, K, buf['out'], BSM, BSN, BSK, NW, NS, WPE, cache)


# B preshuffle view cache (zero-copy reshape, not computation cache)
_b_cache = {}

def custom_kernel(data: input_t) -> output_t:
    A = data[0]
    B_shuffle = data[3]
    B_scale_sh = data[4]
    M, K = A.shape
    N = data[2].shape[0]

    # Cache B preshuffle format (view+reshape only, not GEMM result)
    b_key = (B_shuffle.data_ptr(), N)
    if b_key not in _b_cache:
        B_pre = B_shuffle.view(torch.uint8).reshape(N // 16, (K // 2) * 16)
        bs = B_scale_sh.shape
        B_sc = B_scale_sh.view(torch.uint8).reshape(bs[0] // 32, bs[1] * 32)
        _b_cache[b_key] = (B_pre, B_sc)
    B_pre, B_sc = _b_cache[b_key]

    return _fused_preshuffle(A, B_pre, B_sc, M, N, K)
scrolls · 151 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