Skip to content
KernelIndex
Search⌘K

submission 655804

Hamza · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_direct.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-655804?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
10.5µs
#270 of 1143
2026-03-28

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:23fd7af63261fbe7a83b10e722bb7a304b749892a77a83bb340cf6168f865574
license declaredunknown
license concludedunknown
authorsHamza
imported2026-08-15

Techniques

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

split-k_lines = ["cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio"]
stages = 1STAGES = 1
tile-m = 16BLOCK_M = 16 if M > 8 else 8
tile-n = 128BLOCK_N = 128

Kernel source

submission_direct.py223 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

# --- Config injection ---
import os as _os

_KERNEL_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_CSV_PATH = "/tmp/_mxfp4_mm_config.csv"
_CU = 256
_NK_FAMILIES = [
    (2880, 512), (2112, 7168), (4096, 512), (7168, 2048), (3072, 1536),
    (2880, 1536), (4096, 1536), (2112, 512), (2112, 2048),
    (7168, 512), (7168, 1536), (7168, 7168), (3072, 512),
    (3072, 7168), (3072, 2048), (4096, 2048), (4096, 7168),
    (2880, 2048), (2880, 7168),
]
_M_VALUES = [1, 2, 4, 8, 16, 32, 64, 128, 256]
_lines = ["cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio"]
for _n, _k in _NK_FAMILIES:
    for _m in _M_VALUES:
        _tile_num = ((_m + 31) // 32) * ((_n + 127) // 128)
        _cus_per_tile = _CU / max(_tile_num, 1)
        _split = 0
        while _cus_per_tile >= pow(2, _split + 1) and (pow(2, _split + 1) * 128) < 2 * _k:
            _split += 1
        _split = min(_split, 3)
        _lines.append(f"{_CU},{_m},{_n},{_k},21,{_split},1.0,{_KERNEL_32x128},0,0,0.0")
with open(_CSV_PATH, "w") as _f:
    _f.write("\n".join(_lines))
_os.environ["AITER_CONFIG_GEMM_A4W4"] = _CSV_PATH + ":/home/runner/aiter/aiter/configs/a4w4_blockscale_tuned_gemm.csv"
# --- End config injection ---

import torch
import triton
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,
)
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
from task import input_t, output_t

_PRESHUFFLE_CACHE: dict = {}
_OUT_BUF: dict = {}
_YPP_BUF: dict = {}
_CFG_CACHE: dict = {}


def _get_preshuffle_b(data):
    key = data[3].data_ptr()
    if key not in _PRESHUFFLE_CACHE:
        N = data[3].shape[0]
        K_bytes = data[3].shape[1]
        sm, sn = data[4].shape
        N_groups = N // 32
        B_w = data[3].view(torch.uint8).reshape(N // 16, K_bytes * 16)
        B_s = data[4].view(torch.uint8).reshape(sm // 32, sn * 32)[:N_groups].contiguous()
        _PRESHUFFLE_CACHE[key] = (B_w, B_s)
    return _PRESHUFFLE_CACHE[key]


def _get_cfg(M: int, N: int, K_real: int):
    key = (M, N, K_real)
    if key in _CFG_CACHE:
        return _CFG_CACHE[key]

    K = K_real // 2  # Internal K dimension

    if M < 32:
        BLOCK_M = 16 if M > 8 else 8
        BLOCK_N = 128
        KSPLIT = 1
        STAGES = 1
        if K_real >= 4096:
            KSPLIT = 7
            STAGES = 2
        elif K_real >= 2048:
            KSPLIT = 2
            STAGES = 2
        elif K_real >= 1536:
            KSPLIT = 3
        # Use BLOCK_N=64 when CU utilization is low
        tiles_128 = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + 127) // 128)
        if tiles_128 * KSPLIT < (_CU * 3) // 4:
            BLOCK_N = 64
        cfg = {
            "BLOCK_SIZE_M": BLOCK_M, "BLOCK_SIZE_N": BLOCK_N, "BLOCK_SIZE_K": 512,
            "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": STAGES,
            "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
            "cache_modifier": ".cg", "NUM_KSPLIT": KSPLIT,
        }
    else:
        tiles = ((M + 15) // 16) * ((N + 127) // 128)
        BLOCK_N = 128
        KSPLIT = 1
        STAGES = 2
        # Use KSPLIT to boost CU utilization for shapes with few tiles
        if tiles < _CU // 2 and K_real > 512:
            if K_real >= 4096:
                if tiles * 2 >= _CU:
                    KSPLIT = 2
                else:
                    KSPLIT = 7
            elif K_real >= 2048:
                KSPLIT = 2
            elif K_real >= 1536:
                KSPLIT = 3
                STAGES = 1
        # Use BLOCK_N=64 when CU utilization is low
        if tiles * KSPLIT < (_CU * 3) // 4:
            BLOCK_N = 64
        wgs = ((M + 15) // 16) * ((N + BLOCK_N - 1) // BLOCK_N) * KSPLIT
        cfg = {
            "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": BLOCK_N, "BLOCK_SIZE_K": 512,
            "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": STAGES,
            "waves_per_eu": 2 if wgs > _CU else 1, "matrix_instr_nonkdim": 16,
            "cache_modifier": ".cg", "NUM_KSPLIT": KSPLIT,
        }

    # Apply get_splitk to adjust KSPLIT
    if cfg["NUM_KSPLIT"] > 1:
        SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT = get_splitk(
            K, cfg["BLOCK_SIZE_K"], cfg["NUM_KSPLIT"]
        )
        cfg["SPLITK_BLOCK_SIZE"] = SPLITK_BLOCK_SIZE
        cfg["BLOCK_SIZE_K"] = BLOCK_SIZE_K
        cfg["NUM_KSPLIT"] = NUM_KSPLIT

    # Handle BLOCK_K >= 2*K edge case
    if cfg["BLOCK_SIZE_K"] >= 2 * K:
        cfg["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * K)
        cfg["SPLITK_BLOCK_SIZE"] = 2 * K
        cfg["NUM_KSPLIT"] = 1
    cfg["BLOCK_SIZE_N"] = max(cfg["BLOCK_SIZE_N"], 32)

    if cfg["NUM_KSPLIT"] == 1:
        cfg["SPLITK_BLOCK_SIZE"] = 2 * K

    # Pre-compute reduce kernel params
    actual_ksplit = None
    nk_pow2 = None
    if cfg["NUM_KSPLIT"] > 1:
        actual_ksplit = triton.cdiv(K, cfg["SPLITK_BLOCK_SIZE"] // 2)
        nk_pow2 = triton.next_power_of_2(cfg["NUM_KSPLIT"])

    # Pre-compute grid
    grid_main = (
        cfg["NUM_KSPLIT"]
        * triton.cdiv(M, cfg["BLOCK_SIZE_M"])
        * triton.cdiv(N, cfg["BLOCK_SIZE_N"]),
    )
    grid_reduce = None
    if cfg["NUM_KSPLIT"] > 1:
        grid_reduce = (triton.cdiv(M, 16), triton.cdiv(N, 64))

    result = (cfg, actual_ksplit, nk_pow2, grid_main, grid_reduce)
    _CFG_CACHE[key] = result
    return result


def custom_kernel(data: input_t) -> output_t:
    A = data[0]
    if not A.is_contiguous():
        A = A.contiguous()

    shape_prefix = tuple(A.shape[:-1])
    A_2d = A.view(-1, A.shape[-1])
    M = A_2d.shape[0]
    N = data[3].shape[0]
    K_bytes = data[3].shape[1]
    K_real = K_bytes * 2
    K = K_real // 2

    cfg, actual_ksplit, nk_pow2, grid_main, grid_reduce = _get_cfg(M, N, K_real)

    dev = A.device
    okey = (dev.index, M, N)
    if okey not in _OUT_BUF:
        _OUT_BUF[okey] = torch.empty((M, N), dtype=torch.bfloat16, device=dev)
    y = _OUT_BUF[okey]

    # Pre-allocated y_pp for KSPLIT > 1
    if cfg["NUM_KSPLIT"] > 1:
        ppkey = (dev.index, nk_pow2, M, N)
        if ppkey not in _YPP_BUF:
            _YPP_BUF[ppkey] = torch.empty(
                (nk_pow2, M, N), dtype=torch.float32, device=dev
            )
        y_pp = _YPP_BUF[ppkey]
    else:
        y_pp = None

    B_w, B_s = _get_preshuffle_b(data)

    # Launch main GEMM kernel
    _gemm_a16wfp4_preshuffle_kernel[grid_main](
        A_2d, B_w,
        y if y_pp is None else y_pp,
        B_s,
        M, N, K,
        A_2d.stride(0), A_2d.stride(1),
        B_w.stride(0), B_w.stride(1),
        0 if y_pp is None else y_pp.stride(0),
        y.stride(0) if y_pp is None else y_pp.stride(1),
        y.stride(1) if y_pp is None else y_pp.stride(2),
        B_s.stride(0), B_s.stride(1),
        PREQUANT=True,
        **cfg,
    )

    # Reduce if KSPLIT > 1
    if y_pp is not None:
        _gemm_afp4wfp4_reduce_kernel[grid_reduce](
            y_pp, y, M, N,
            y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
            y.stride(0), y.stride(1),
            16, 64,
            actual_ksplit, nk_pow2,
        )

    return y.view(*shape_prefix, N)
scrolls · 223 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 654326.

⋯ 68 unchanged lines
if M < 32:
BLOCK_M = 16 if M > 8 else 8
+ BLOCK_N = 128
KSPLIT = 1
STAGES = 1
if K_real >= 4096:
⋯ 3 unchanged lines
KSPLIT = 2
STAGES = 2
elif K_real >= 1536:
- KSPLIT = 4
+ KSPLIT = 3
+ # Use BLOCK_N=64 when CU utilization is low
+ tiles_128 = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + 127) // 128)
+ if tiles_128 * KSPLIT < (_CU * 3) // 4:
+ BLOCK_N = 64
cfg = {
- "BLOCK_SIZE_M": BLOCK_M, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
+ "BLOCK_SIZE_M": BLOCK_M, "BLOCK_SIZE_N": BLOCK_N, "BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": STAGES,
"waves_per_eu": 1, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": KSPLIT,
}
else:
tiles = ((M + 15) // 16) * ((N + 127) // 128)
+ BLOCK_N = 128
KSPLIT = 1
STAGES = 2
# Use KSPLIT to boost CU utilization for shapes with few tiles
if tiles < _CU // 2 and K_real > 512:
if K_real >= 4096:
- # K=3584: KSPLIT=7→2 iters, KSPLIT=2→7 iters
if tiles * 2 >= _CU:
KSPLIT = 2
else:
KSPLIT = 7
elif K_real >= 2048:
- # K=1024: KSPLIT=2→2 iters with stages=2
KSPLIT = 2
elif K_real >= 1536:
- # K=768: KSPLIT=3→1 iter (no pipeline benefit)
KSPLIT = 3
STAGES = 1
- wgs = tiles * KSPLIT
+ # Use BLOCK_N=64 when CU utilization is low
+ if tiles * KSPLIT < (_CU * 3) // 4:
+ BLOCK_N = 64
+ wgs = ((M + 15) // 16) * ((N + BLOCK_N - 1) // BLOCK_N) * KSPLIT
cfg = {
- "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
+ "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": BLOCK_N, "BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": STAGES,
"waves_per_eu": 2 if wgs > _CU else 1, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": KSPLIT,
⋯ 63 unchanged lines
# Pre-allocated y_pp for KSPLIT > 1
if cfg["NUM_KSPLIT"] > 1:
- ppkey = (dev.index, cfg["NUM_KSPLIT"], M, N)
+ ppkey = (dev.index, nk_pow2, M, N)
if ppkey not in _YPP_BUF:
_YPP_BUF[ppkey] = torch.empty(
- (cfg["NUM_KSPLIT"], M, N), dtype=torch.float32, device=dev
+ (nk_pow2, M, N), dtype=torch.float32, device=dev
)
y_pp = _YPP_BUF[ppkey]
else:
scrolls · 69 diff lines total

Best evidence level for this revision: reported

JSON