Skip to content
KernelIndex
Search⌘K

submission 585029

zhaohb · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_opt.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-585029?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.4µs
#791 of 1143
2026-03-18

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:99e4949dbfbc42aab141442194727f8ea9828cb5011b93992101b013dd077334
license declaredunknown
license concludedunknown
authorszhaohb
imported2026-08-26

Techniques

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

fp4Optimized FP4 quant + FP4 GEMM path for MI355X.
num-warps = 1num_warps = 1
split-ksplit_k = 0
stages = 1num_stages=1,

Kernel source

submission_opt.py223 lines
"""
Optimized FP4 quant + FP4 GEMM path for MI355X.

Optimizations:
1. Move imports/constants to module scope to trim Python overhead.
2. Reuse quant/output buffers across repeated calls with the same shapes.
3. Bypass `aiter.gemm_a4w4()`'s per-call output allocation when internals are available.
4. Fall back to the public aiter path if any internal fast path is unavailable.
"""
from task import input_t, output_t

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

try:
    import triton
    from aiter.ops.triton._triton_kernels.quant.quant import _dynamic_mxfp4_quant_kernel

    _HAS_FAST_QUANT = True
except Exception:
    triton = None
    _dynamic_mxfp4_quant_kernel = None
    _HAS_FAST_QUANT = False

try:
    from aiter.ops.gemm_op_a4w4 import (
        gemm_a4w4_asm,
        gemm_a4w4_blockscale,
        get_GEMM_config,
    )

    _HAS_FAST_GEMM = True
except Exception:
    gemm_a4w4_asm = None
    gemm_a4w4_blockscale = None
    get_GEMM_config = None
    _HAS_FAST_GEMM = False


BF16 = dtypes.bf16
FP4X2 = dtypes.fp4x2
FP8_E8M0 = dtypes.fp8_e8m0
_PUBLIC_GEMM_A4W4 = aiter.gemm_a4w4

_FAST_PATH_ENABLED = _HAS_FAST_QUANT and _HAS_FAST_GEMM
_QUANT_CACHE: dict[tuple[object, int, int], tuple[torch.Tensor, torch.Tensor, torch.Tensor]] = {}
_SCALE_SHUFFLE_CACHE: dict[
    tuple[object, int, int],
    tuple[torch.Tensor, torch.Tensor, torch.Tensor],
] = {}
_OUT_CACHE: dict[tuple[object, int, int], torch.Tensor] = {}


def _device_key(device: torch.device) -> tuple[str, int | None]:
    return (device.type, device.index)


def _quant_mxfp4_cached(x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    m, k = x.shape
    cache_key = (_device_key(x.device), m, k)
    cached = _QUANT_CACHE.get(cache_key)
    if cached is None:
        x_fp4_u8 = torch.empty((m, k // 2), dtype=torch.uint8, device=x.device)
        scale_u8 = torch.empty((m, (k + 31) // 32), dtype=torch.uint8, device=x.device)
        x_fp4 = x_fp4_u8.view(FP4X2)
        _QUANT_CACHE[cache_key] = (x_fp4_u8, scale_u8, x_fp4)
    else:
        x_fp4_u8, scale_u8, x_fp4 = cached

    if m <= 32:
        num_iter = 1
        block_size_m = triton.next_power_of_2(m)
        block_size_n = 32
        num_warps = 1
        num_stages_cfg = 1
    else:
        num_iter = 4
        block_size_m = 64
        block_size_n = 64
        num_warps = 4
        num_stages_cfg = 2

    if k <= 16384:
        block_size_m = 32
        block_size_n = 128

    if k <= 1024:
        num_iter = 1
        num_stages_cfg = 1
        num_warps = 4
        block_size_n = max(32, min(256, triton.next_power_of_2(k)))
        block_size_m = min(8, triton.next_power_of_2(m))

    grid = (
        triton.cdiv(m, block_size_m),
        triton.cdiv(k, block_size_n * num_iter),
    )
    _dynamic_mxfp4_quant_kernel[grid](
        x,
        x_fp4_u8,
        scale_u8,
        *x.stride(),
        *x_fp4_u8.stride(),
        *scale_u8.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_cfg,
        num_warps=num_warps,
        waves_per_eu=0,
        num_stages=1,
    )
    return x_fp4, scale_u8


def _shuffle_e8m0_cached(scale_u8: torch.Tensor) -> torch.Tensor:
    m, n = scale_u8.shape
    padded_m = (m + 255) // 256 * 256
    padded_n = (n + 7) // 8 * 8
    cache_key = (_device_key(scale_u8.device), padded_m, padded_n)
    cached = _SCALE_SHUFFLE_CACHE.get(cache_key)
    if cached is None:
        scale_pad = torch.empty((padded_m, padded_n), dtype=torch.uint8, device=scale_u8.device)
        scale_sh = torch.empty((padded_m, padded_n), dtype=torch.uint8, device=scale_u8.device)
        scale_sh_view = scale_sh.view(padded_m // 32, padded_n // 8, 4, 16, 2, 2)
        _SCALE_SHUFFLE_CACHE[cache_key] = (scale_pad, scale_sh, scale_sh_view)
    else:
        scale_pad, scale_sh, scale_sh_view = cached

    scale_pad[:m, :n] = scale_u8
    scale_sh_view.copy_(
        scale_pad.view(padded_m // 32, 2, 16, padded_n // 8, 2, 4).permute(0, 3, 5, 2, 4, 1)
    )
    return scale_sh.view(FP8_E8M0)


def _get_out_buffer(device: torch.device, m: int, n: int) -> torch.Tensor:
    padded_m = (m + 31) // 32 * 32
    cache_key = (_device_key(device), padded_m, n)
    out = _OUT_CACHE.get(cache_key)
    if out is None:
        out = torch.empty((padded_m, n), dtype=BF16, device=device)
        _OUT_CACHE[cache_key] = out
    return out


def _gemm_cached(
    a_q: torch.Tensor,
    b_shuffle: torch.Tensor,
    a_scale_sh: torch.Tensor,
    b_scale_sh: torch.Tensor,
) -> torch.Tensor:
    m = a_q.shape[0]
    n = b_shuffle.shape[0]
    k = a_q.shape[1] * 2
    out = _get_out_buffer(a_q.device, m, n)

    ck_config = get_GEMM_config(m, n, k)
    split_k = 0
    kernel_name = ""
    if ck_config is not None:
        split_k = ck_config.get("splitK", None)
        kernel_name = ck_config["kernelName"]

    if ck_config is not None and "_ZN" not in kernel_name:
        split_k = 0 if split_k is None else split_k
        gemm_a4w4_blockscale(a_q, b_shuffle, a_scale_sh, b_scale_sh, out, splitK=split_k)
    else:
        gemm_a4w4_asm(
            a_q,
            b_shuffle,
            a_scale_sh,
            b_scale_sh,
            out,
            kernel_name,
            None,
            1.0,
            0.0,
            True,
            log2_k_split=split_k,
        )

    return out[:m]


def _public_path(a: torch.Tensor, b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor) -> torch.Tensor:
    x_fp4, scale_e8m0 = _public_dynamic_mxfp4_quant(a)
    a_q = x_fp4.view(FP4X2)
    a_scale_sh = _public_e8m0_shuffle(scale_e8m0).view(FP8_E8M0)
    return _PUBLIC_GEMM_A4W4(
        a_q,
        b_shuffle,
        a_scale_sh,
        b_scale_sh,
        dtype=BF16,
        bpreshuffle=True,
    )


def custom_kernel(data: input_t) -> output_t:
    global _FAST_PATH_ENABLED

    a, _, _, b_shuffle, b_scale_sh = data
    if not a.is_contiguous():
        a = a.contiguous()

    if _FAST_PATH_ENABLED:
        try:
            a_q, a_scale = _quant_mxfp4_cached(a)
            a_scale_sh = _shuffle_e8m0_cached(a_scale)
            return _gemm_cached(a_q, b_shuffle, a_scale_sh, b_scale_sh)
        except Exception:
            _FAST_PATH_ENABLED = False

    return _public_path(a, b_shuffle, b_scale_sh)
scrolls · 223 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