Skip to content
KernelIndex
Search⌘K

submission 662095

Hamza · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:356286fc10c0da3f20a7b174a7ccebaa3dc430af47f328f49326785876aaa343
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 = 2STAGES = 2
tile-k = 256BLOCK_K = 256 if K_real <= KSPLIT * 512 else 512
tile-m = 8BLOCK_M = 8
tile-n = 128BLOCK_N = 128

Kernel source

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

# --- Config injection (prevents extra module_gemm_common build) ---
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 task import input_t, output_t


def _get_splitk(K: int, BLOCK_SIZE_K: int, NUM_KSPLIT: int):
    """Adjust KSPLIT/BLOCK_K for EVEN_K alignment (inlined from aiter)."""
    SPLITK_BLOCK_SIZE = (
        triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K
    )
    while NUM_KSPLIT > 1 and BLOCK_SIZE_K > 16:
        if (
            K % (SPLITK_BLOCK_SIZE // 2) == 0
            and SPLITK_BLOCK_SIZE % BLOCK_SIZE_K == 0
            and K % (BLOCK_SIZE_K // 2) == 0
        ):
            break
        elif K % (SPLITK_BLOCK_SIZE // 2) != 0 and NUM_KSPLIT > 1:
            NUM_KSPLIT = NUM_KSPLIT // 2
        elif SPLITK_BLOCK_SIZE % BLOCK_SIZE_K != 0:
            if NUM_KSPLIT > 1:
                NUM_KSPLIT = NUM_KSPLIT // 2
            elif BLOCK_SIZE_K > 16:
                BLOCK_SIZE_K = BLOCK_SIZE_K // 2
        elif K % (BLOCK_SIZE_K // 2) != 0 and BLOCK_SIZE_K > 16:
            BLOCK_SIZE_K = BLOCK_SIZE_K // 2
        else:
            break
        SPLITK_BLOCK_SIZE = (
            triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K
        )
    return SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT


_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 = 8
        BLOCK_N = 128
        KSPLIT = 1
        STAGES = 2
        if K_real >= 4096:
            KSPLIT = 7
        elif K_real >= 2048:
            KSPLIT = 4
        elif K_real >= 1536:
            KSPLIT = 3
        # Use BLOCK_K=256 when each K-split has ≤1 iter with BK=512 → enables pipeline
        BLOCK_K = 256 if K_real <= KSPLIT * 512 else 512
        # 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
        wgs = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + BLOCK_N - 1) // BLOCK_N) * KSPLIT
        cfg = {
            "BLOCK_SIZE_M": BLOCK_M, "BLOCK_SIZE_N": BLOCK_N, "BLOCK_SIZE_K": BLOCK_K,
            "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,
        }
    else:
        # Use BLOCK_M=8 for M<=64 when CU utilization with BM=16 is low
        BLOCK_M = 16
        if M <= 128:
            tiles_bm16 = ((M + 15) // 16) * ((N + 127) // 128)
            if tiles_bm16 < (_CU * 3) // 4:
                BLOCK_M = 8
        tiles = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + 127) // 128)
        BLOCK_N = 128
        KSPLIT = 1
        STAGES = 2
        # Use KSPLIT to boost CU utilization for shapes with few tiles
        if K_real >= 7168 and _CU // 2 <= tiles <= _CU:
            # K=7168 has 14 K-iters: KSPLIT=2 halves to 7 with small reduce overhead
            KSPLIT = 2
        elif 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
        # KSPLIT=2 to reduce wave tail for 1.x-wave shapes with K=2048
        if KSPLIT == 1 and _CU < tiles <= 2 * _CU and K_real == 2048:
            KSPLIT = 2
        # Use BLOCK_K=256 when each K-split has ≤1 iter with BK=512 → enables pipeline
        BLOCK_K = 256 if K_real <= KSPLIT * 512 else 512
        # Use BLOCK_N=64 when CU utilization is low
        if tiles * KSPLIT < (_CU * 3) // 4:
            BLOCK_N = 64
        wgs = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + BLOCK_N - 1) // BLOCK_N) * KSPLIT
        cfg = {
            "BLOCK_SIZE_M": BLOCK_M, "BLOCK_SIZE_N": BLOCK_N, "BLOCK_SIZE_K": BLOCK_K,
            "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
    num_m_tiles = triton.cdiv(M, cfg["BLOCK_SIZE_M"])
    num_n_tiles = triton.cdiv(N, cfg["BLOCK_SIZE_N"])
    total_tiles = num_m_tiles * num_n_tiles
    grid_main = (cfg["NUM_KSPLIT"] * total_tiles,)
    grid_reduce = None
    if cfg["NUM_KSPLIT"] > 1:
        grid_reduce = (triton.cdiv(M, 16), triton.cdiv(N, 16))

    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]

    B_w, B_s = _get_preshuffle_b(data)

    # 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

    # 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, 16,
            actual_ksplit, nk_pow2,
        )

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

⋯ 99 unchanged lines
BLOCK_M = 8
BLOCK_N = 128
KSPLIT = 1
- STAGES = 1
+ STAGES = 2
if K_real >= 4096:
KSPLIT = 7
- STAGES = 2
elif K_real >= 2048:
- KSPLIT = 2
- STAGES = 2
+ KSPLIT = 4
elif K_real >= 1536:
KSPLIT = 3
- STAGES = 2
+ # Use BLOCK_K=256 when each K-split has ≤1 iter with BK=512 → enables pipeline
+ BLOCK_K = 256 if K_real <= KSPLIT * 512 else 512
# 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
wgs = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + BLOCK_N - 1) // BLOCK_N) * KSPLIT
cfg = {
- "BLOCK_SIZE_M": BLOCK_M, "BLOCK_SIZE_N": BLOCK_N, "BLOCK_SIZE_K": 512,
+ "BLOCK_SIZE_M": BLOCK_M, "BLOCK_SIZE_N": BLOCK_N, "BLOCK_SIZE_K": BLOCK_K,
"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,
}
else:
- tiles = ((M + 15) // 16) * ((N + 127) // 128)
+ # Use BLOCK_M=8 for M<=64 when CU utilization with BM=16 is low
+ BLOCK_M = 16
+ if M <= 128:
+ tiles_bm16 = ((M + 15) // 16) * ((N + 127) // 128)
+ if tiles_bm16 < (_CU * 3) // 4:
+ BLOCK_M = 8
+ tiles = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + 127) // 128)
BLOCK_N = 128
KSPLIT = 1
- STAGES = 1 if K_real <= 512 else 2 # 1 K-iter → no pipeline benefit
+ STAGES = 2
# Use KSPLIT to boost CU utilization for shapes with few tiles
- if K_real >= 7168 and _CU // 2 <= tiles < _CU:
+ if K_real >= 7168 and _CU // 2 <= tiles <= _CU:
# K=7168 has 14 K-iters: KSPLIT=2 halves to 7 with small reduce overhead
KSPLIT = 2
elif tiles < _CU // 2 and K_real > 512:
⋯ 6 unchanged lines
KSPLIT = 2
elif K_real >= 1536:
KSPLIT = 3
- STAGES = 2
+ # KSPLIT=2 to reduce wave tail for 1.x-wave shapes with K=2048
+ if KSPLIT == 1 and _CU < tiles <= 2 * _CU and K_real == 2048:
+ KSPLIT = 2
+ # Use BLOCK_K=256 when each K-split has ≤1 iter with BK=512 → enables pipeline
+ BLOCK_K = 256 if K_real <= KSPLIT * 512 else 512
# 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
+ wgs = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + BLOCK_N - 1) // BLOCK_N) * KSPLIT
cfg = {
- "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": BLOCK_N, "BLOCK_SIZE_K": 512,
+ "BLOCK_SIZE_M": BLOCK_M, "BLOCK_SIZE_N": BLOCK_N, "BLOCK_SIZE_K": BLOCK_K,
"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,
scrolls · 70 diff lines total

Best evidence level for this revision: reported

JSON