Skip to content
KernelIndex
Search⌘K

submission 692216

ZainHaider20 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-692216?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
20.6µs
#728 of 1143
2026-04-01

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:d1e4f35889ccc6a6a4e2302ff4d045c64c616783aad55c9b46e9f7b27d1d1e34
license declaredunknown
license concludedunknown
authorsZainHaider20
imported2026-08-26

Techniques

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

split-k_CUSTOM_CSV = """cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio

Kernel source

submission.py190 lines
"""
Use injected-config wrapper GEMM only for the m=4 and m=16 public benchmark
shapes, while keeping the stable direct ASM path for the two m=32 shapes and
the large shapes.
"""
from __future__ import annotations

import os
from pathlib import Path

from task import input_t, output_t


_CUSTOM_CSV = """cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio
256,4,2880,512,29,0,0.0,_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E,0.0,0.0,0.0
256,16,2112,7168,29,1,0.0,_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E,0.0,0.0,0.0
"""

_CONFIG_PATH = Path("/tmp/aiter_mxfp4_mm_cfg_m4_m16_only.csv")
if not _CONFIG_PATH.exists():
    _CONFIG_PATH.write_text(_CUSTOM_CSV, encoding="utf-8")
os.environ["AITER_CONFIG_GEMM_A4W4"] = str(_CONFIG_PATH)

import aiter
import torch
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle


_WRAPPER_SHAPES = {
    (4, 2880, 512),
    (8, 2112, 7168),
    (16, 2112, 7168),
}
_ASM_KERNEL_MAP: dict[tuple[int, int, int], tuple[str, int]] = {
    (32, 4096, 512): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        0,
    ),
    (32, 2880, 512): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        0,
    ),
    (64, 7168, 2048): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        0,
    ),
    (256, 3072, 1536): (
        "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
        0,
    ),
}
_EXACT_CACHED_QUANT_SHAPES = {
    (64, 7168, 2048),
    (256, 3072, 1536),
}
_EXACT_QUANT_CONFIGS = {
    (64, 2048): (4, 32, 128, 4, 2),
    (64, 1536): (4, 32, 128, 4, 2),
    (256, 1536): (4, 32, 128, 4, 2),
}
_OUT_CACHE: dict[tuple[int, int, torch.device], torch.Tensor] = {}
_QUANT_CACHE: dict[tuple[int, int, torch.device], tuple[torch.Tensor, torch.Tensor]] = {}
_TRITON_BITS = None


def _get_triton_bits():
    global _TRITON_BITS
    if _TRITON_BITS is not None:
        return _TRITON_BITS
    import triton
    from aiter.ops.triton._triton_kernels.quant.quant import _dynamic_mxfp4_quant_kernel

    _TRITON_BITS = (triton, _dynamic_mxfp4_quant_kernel)
    return _TRITON_BITS


@torch.inference_mode()
def _get_quant_buffers(m: int, k: int, device: torch.device):
    key = (m, k, device)
    cached = _QUANT_CACHE.get(key)
    if cached is not None:
        return cached
    x_fp4 = torch.empty((m, k // 2), dtype=torch.uint8, device=device)
    blockscale_e8m0 = torch.empty((m, k // 32), dtype=torch.uint8, device=device)
    _QUANT_CACHE[key] = (x_fp4, blockscale_e8m0)
    return x_fp4, blockscale_e8m0


@torch.inference_mode()
def _quant_mxfp4(x: torch.Tensor):
    x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
    return x_fp4.view(dtypes.fp4x2), e8m0_shuffle(bs_e8m0).view(dtypes.fp8_e8m0)


_compiled_quant_mxfp4 = torch.compile(_quant_mxfp4)


@torch.inference_mode()
def _quant_mxfp4_cached_exact(x: torch.Tensor):
    triton, kernel = _get_triton_bits()
    m, k = x.shape
    x_fp4, blockscale_e8m0 = _get_quant_buffers(m, k, x.device)
    num_iter, block_size_m, block_size_n, num_warps, num_stages = _EXACT_QUANT_CONFIGS[(m, k)]
    grid = (
        triton.cdiv(m, block_size_m),
        triton.cdiv(k, block_size_n * num_iter),
    )
    kernel[grid](
        x,
        x_fp4,
        blockscale_e8m0,
        *x.stride(),
        *x_fp4.stride(),
        *blockscale_e8m0.stride(),
        M=m,
        N=k,
        MXFP4_QUANT_BLOCK_SIZE=32,
        SCALING_MODE=0,
        NUM_ITER=num_iter,
        BLOCK_SIZE_M=block_size_m,
        BLOCK_SIZE_N=block_size_n,
        NUM_STAGES=num_stages,
        num_stages=num_stages,
        num_warps=num_warps,
        waves_per_eu=0,
    )
    return x_fp4.view(dtypes.fp4x2), e8m0_shuffle(blockscale_e8m0).view(dtypes.fp8_e8m0)


@torch.inference_mode()
def _get_out(m: int, n: int, device: torch.device) -> torch.Tensor:
    key = (m, n, device)
    out = _OUT_CACHE.get(key)
    if out is None:
        out = torch.empty(((m + 31) // 32 * 32, n), dtype=dtypes.bf16, device=device)
        _OUT_CACHE[key] = out
    return out


@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
    A, _, _, B_shuffle, B_scale_sh = data
    if not A.is_contiguous():
        A = A.contiguous()
    m, k = A.shape
    n = B_shuffle.shape[0]
    shape = (m, n, k)

    if shape in _WRAPPER_SHAPES:
        A_q, A_scale_sh = _compiled_quant_mxfp4(A)
        return aiter.gemm_a4w4(
            A_q,
            B_shuffle,
            A_scale_sh,
            B_scale_sh,
            dtype=dtypes.bf16,
            bpreshuffle=True,
        )

    if shape in _EXACT_CACHED_QUANT_SHAPES:
        A_q, A_scale_sh = _quant_mxfp4_cached_exact(A)
    else:
        A_q, A_scale_sh = _compiled_quant_mxfp4(A)

    if shape not in _ASM_KERNEL_MAP:
        return aiter.gemm_a4w4(
            A_q,
            B_shuffle,
            A_scale_sh,
            B_scale_sh,
            dtype=dtypes.bf16,
            bpreshuffle=True,
        )

    kernel_name, split_k = _ASM_KERNEL_MAP[shape]
    out = _get_out(m, n, A.device)
    aiter.gemm_a4w4_asm(
        A_q.view(m, k // 2),
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        out,
        kernel_name,
        bpreshuffle=True,
        log2_k_split=split_k,
    )
    return out[:m]
scrolls · 190 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