Skip to content
KernelIndex
Search⌘K

submission 651591

Hamza · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:13e97f76437f368035c95882a769d61541e6226a15a760a7976de8fe59105448
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_tuned.py120 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
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
from task import input_t, output_t

_PRESHUFFLE_CACHE: dict = {}
_OUT_BUF: 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_preshuffle_config(M: int, N: int, K_real: int) -> dict:
    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
        return {
            "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,
            "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,
        }


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

    okey = (A.device.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]

    # Preshuffle for all M (non-preshuffle has ranked correctness issues)
    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)
scrolls · 120 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 526320.

- """
- FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.
- Quant logic follows aiter op_tests/test_gemm_a4w4.py (get_triton_quant(QuantType.per_1x32)).
- """
- import aiter
- from aiter import QuantType, dtypes
+ #!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
+ from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
from task import input_t, output_t
- QUANT_FUNC = aiter.get_triton_quant(QuantType.per_1x32)
- GEMM_A4W4 = aiter.gemm_a4w4
- OUTPUT_DTYPE = dtypes.bf16
+ _PRESHUFFLE_CACHE: dict = {}
+ _OUT_BUF: 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_preshuffle_config(M: int, N: int, K_real: int) -> dict:
+ 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
+ return {
+ "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,
+ "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,
+ }
+
+
def custom_kernel(data: input_t) -> output_t:
- """
- Reference: MXFP4 per-1x32 quant on A; B_shuffle, B_scale_sh from generate_input.
- gemm_a4w4 with bpreshuffle=True.
- """
- A, _, _, B_shuffle, B_scale_sh = data
- A_q, A_scale_sh = QUANT_FUNC(A, shuffle=True)
+ A = data[0]
+ if not A.is_contiguous():
+ A = A.contiguous()
- return GEMM_A4W4(
- A_q,
- B_shuffle,
- A_scale_sh,
- B_scale_sh,
- dtype=OUTPUT_DTYPE,
- bpreshuffle=True,
- )
+ 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
+
+ okey = (A.device.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]
+
+ # Preshuffle for all M (non-preshuffle has ranked correctness issues)
+ 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)
scrolls · 142 diff lines total

Best evidence level for this revision: reported

JSON