Skip to content
KernelIndex
Search⌘K

submission 652726

Hamza · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:cc407d9e6ac9fe21108262e62a6065998e494af8a1a0e77f333330d735d65173
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"]
tile-m = 16BLOCK_M = 16 if M > 8 else 8

Kernel source

submission_direct.py193 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
        KSPLIT = 1
        if K_real >= 4096:
            KSPLIT = 14
        elif K_real >= 1536:
            KSPLIT = 4
        cfg = {
            "BLOCK_SIZE_M": BLOCK_M, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
            "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 1,
            "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
            "cache_modifier": ".cg", "NUM_KSPLIT": KSPLIT,
        }
    else:
        cfg = {
            "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
            "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
            "waves_per_eu": 2, "matrix_instr_nonkdim": 16,
            "cache_modifier": ".cg", "NUM_KSPLIT": 1,
        }

    # 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, cfg["NUM_KSPLIT"], M, N)
        if ppkey not in _YPP_BUF:
            _YPP_BUF[ppkey] = torch.empty(
                (cfg["NUM_KSPLIT"], 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 · 193 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 651591.

⋯ 30 unchanged lines
# --- End config injection ---
import torch
- from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
+ 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):
⋯ 9 unchanged lines
return _PRESHUFFLE_CACHE[key]
- def _get_preshuffle_config(M: int, N: int, K_real: int) -> dict:
+ 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
KSPLIT = 1
⋯ 1 unchanged lines
KSPLIT = 14
elif K_real >= 1536:
KSPLIT = 4
- return {
+ cfg = {
"BLOCK_SIZE_M": BLOCK_M, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 1,
"waves_per_eu": 1, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": KSPLIT,
}
- elif M <= 32:
- # M=32: BLOCK_M=16 doubles tile count, BLOCK_K=512 (low register pressure)
- KSPLIT = 1
- if K_real >= 4096:
- KSPLIT = 14
- elif K_real >= 1536:
- KSPLIT = 4
- return {
- "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
- "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 1,
- "waves_per_eu": 2, "matrix_instr_nonkdim": 16,
- "cache_modifier": ".cg", "NUM_KSPLIT": KSPLIT,
- }
else:
- # M>=64: BLOCK_M=32, BLOCK_K=256 (proven, avoids register spilling)
- if K_real >= 4096:
- target_ksplit = 4
- elif K_real >= 1536:
- target_ksplit = 2
- else:
- target_ksplit = 1
- return {
- "BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256,
+ cfg = {
+ "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
"waves_per_eu": 2, "matrix_instr_nonkdim": 16,
- "cache_modifier": ".cg", "NUM_KSPLIT": target_ksplit,
+ "cache_modifier": ".cg", "NUM_KSPLIT": 1,
}
+ # 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():
⋯ 5 unchanged lines
N = data[3].shape[0]
K_bytes = data[3].shape[1]
K_real = K_bytes * 2
+ K = K_real // 2
- okey = (A.device.index, M, N)
+ 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=A.device)
- out = _OUT_BUF[okey]
+ _OUT_BUF[okey] = torch.empty((M, N), dtype=torch.bfloat16, device=dev)
+ y = _OUT_BUF[okey]
- # Preshuffle for all M (non-preshuffle has ranked correctness issues)
+ # Pre-allocated y_pp for KSPLIT > 1
+ if cfg["NUM_KSPLIT"] > 1:
+ ppkey = (dev.index, cfg["NUM_KSPLIT"], M, N)
+ if ppkey not in _YPP_BUF:
+ _YPP_BUF[ppkey] = torch.empty(
+ (cfg["NUM_KSPLIT"], M, N), dtype=torch.float32, device=dev
+ )
+ y_pp = _YPP_BUF[ppkey]
+ else:
+ y_pp = None
+
B_w, B_s = _get_preshuffle_b(data)
- config = _get_preshuffle_config(M, N, K_real)
- result = gemm_a16wfp4_preshuffle(A_2d, B_w, B_s, y=out, config=config)
- return result.view(*shape_prefix, N)
+ # 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 · 184 diff lines total

Best evidence level for this revision: reported

JSON