Skip to content
KernelIndex
Search⌘K

submission 543606

GRIM5th · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

test10.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-543606?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
23.5µs
#797 of 1143
2026-03-13

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:d2ffd6a542cd8d3b36bf0924350378d1bd2deb6715d1b58f322b58f76925f612
license declaredunknown
license concludedunknown
authorsGRIM5th
imported2026-08-26

Techniques

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

fp4FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.
num-warps = 4NUM_WARPS = 4
stages = 2NUM_STAGES = 2
tile-m = 64BLOCK_SIZE_M = 64
tile-n = 128BLOCK_SIZE_N = 128

Kernel source

test10.py112 lines
# 1.2B *2 ops 
"""
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)).

NOTE: Explicitly uses dynamic_mxfp4_quant from aiter.ops.triton.quant (patched in #975)
      rather than going through aiter.get_triton_quant, which may dispatch to the
      unpatched fp4_utils.py kernel. See ROCm/aiter#974, ROCm/aiter#975.
"""
import torch
import triton
import aiter
from task import input_t, output_t
from utils import make_match_reference
from aiter import QuantType,dtypes
from aiter.ops.shuffle import shuffle_weight
from aiter.ops.triton.quant import dynamic_mxfp4_quant  # #975-patched kernel
from aiter.utility.fp4_utils import e8m0_shuffle
from aiter.ops.triton.quant.quant import (
    _static_per_tensor_quant_fp8_i8_kernel,
    _dynamic_per_tensor_quant_fp8_i8_kernel,
    _dynamic_per_token_quant_fp8_i8_kernel,
    _dynamic_mxfp4_quant_kernel,
    _mxfp4_quant_op,
)
from aiter.ops.triton.utils.logger import AiterTritonLogger
# K must be divisible by 64 (scale group 32 and fp4 pack 2)
SCALE_GROUP_SIZE = 32

def _quant_mxfp4(x, shuffle=True):
    #x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
    x_fp4, bs_e8m0 = custom_mxfp4_quant(x)
    if shuffle:
        bs_e8m0 = e8m0_shuffle(bs_e8m0)
    return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)

def custom_mxfp4_quant(x: torch.Tensor, scaling_mode: str = "even"):
    M, K = x.shape
    MXFP4_QUANT_BLOCK_SIZE = SCALE_GROUP_SIZE

    assert (K // 2) % 2 == 0

    device = x.device

    x_fp4 = torch.empty((M, K // 2), dtype=torch.uint8, device=device)

    num_scale_blocks = (K + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
    blockscale_e8m0 = torch.empty((M, num_scale_blocks), dtype=torch.uint8, device=device)

    # heuristic based on MI355X architecture
    # source 
    # https://chipsandcheese.com/p/amds-cdna-4-architecture-announcement
    if K <= 1024:
        BLOCK_SIZE_M = min(32, triton.next_power_of_2(M))
        BLOCK_SIZE_N = 128
        NUM_WARPS = 4
        NUM_STAGES = 2

    elif K <= 16384:
        BLOCK_SIZE_M = 64
        BLOCK_SIZE_N = 256
        NUM_WARPS = 4
        NUM_STAGES = 2

    else:
        BLOCK_SIZE_M = 64
        BLOCK_SIZE_N = 512
        NUM_WARPS = 8
        NUM_STAGES = 3

    grid = (
        triton.cdiv(M, BLOCK_SIZE_M),
        triton.cdiv(K, BLOCK_SIZE_N),
    )

    _dynamic_mxfp4_quant_kernel[grid](
        x,
        x_fp4,
        blockscale_e8m0,
        *x.stride(),
        *x_fp4.stride(),
        *blockscale_e8m0.stride(),
        M=M,
        N=K,
        MXFP4_QUANT_BLOCK_SIZE=MXFP4_QUANT_BLOCK_SIZE,
        SCALING_MODE=0,
        NUM_ITER=1,
        BLOCK_SIZE_M=BLOCK_SIZE_M,
        BLOCK_SIZE_N=BLOCK_SIZE_N,
        NUM_STAGES=NUM_STAGES,
        num_warps=NUM_WARPS,
        waves_per_eu=0,
        num_stages=NUM_STAGES,
    )

    return x_fp4, blockscale_e8m0
    
def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    A_q, A_scale_sh = _quant_mxfp4(A, shuffle=True)
    out_gemm = aiter.gemm_a4w4(
        A_q,
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )
    return out_gemm


#check_implementation = make_match_reference(custom_kernel, rtol=1e-02, atol=1e-02)
scrolls · 112 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