Skip to content
KernelIndex
Search⌘K

submission 516735

Young Han · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-516735?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
13.3µs
#421 of 1143
2026-03-08

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:95c80eea359cbd5911b78d59c87fe98f5ccdd4635380ddf1377b13a0252905f1
license declaredunknown
license concludedunknown
authorsYoung Han
imported2026-08-26

Techniques

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

split-ksk = config.get("splitK", 0)

Kernel source

submission.py134 lines
import os
os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")

import aiter
import torch
import triton
from aiter import QuantType, dtypes
from aiter.utility import fp4_utils
from task import input_t, output_t

try:
    from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm, gemm_a4w4_blockscale, get_GEMM_config
except Exception:
    gemm_a4w4_asm = None
    gemm_a4w4_blockscale = None
    get_GEMM_config = None

_hip_quant = None
try:
    _hip_quant = aiter.get_hip_quant(QuantType.per_1x32)
except Exception:
    pass

_triton_quant = aiter.get_triton_quant(QuantType.per_1x32)
_quant_kernel = getattr(fp4_utils, "_dynamic_mxfp4_quant_kernel_asm_layout", None)

_ASM_32X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_ASM_64X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E"

_quant_cache: dict = {}
_out_cache: dict = {}
_dispatch_cache: dict = {}


def _quantize_a_hip(a: torch.Tensor):
    return _hip_quant(a, shuffle=True)


def _quantize_a_triton(a: torch.Tensor):
    if _quant_kernel is None:
        return _triton_quant(a, shuffle=True)

    m, k = a.shape
    key = (a.device, m, k)
    buf = _quant_cache.get(key)
    if buf is None:
        scale_n_valid = k // 32
        scale_n_pad = triton.cdiv(scale_n_valid, 8) * 8
        scale_m_pad = triton.cdiv(m, 32) * 32
        q_u8 = torch.empty((m, k // 2), dtype=torch.uint8, device=a.device)
        scale_u8 = torch.empty(
            (triton.cdiv(m, 256) * 256, scale_n_pad),
            dtype=torch.uint8, device=a.device,
        )
        buf = (q_u8, scale_u8, scale_n_valid, scale_n_pad, scale_m_pad)
        _quant_cache[key] = buf

    q_u8, scale_u8, scale_n_valid, scale_n_pad, scale_m_pad = buf
    grid = (triton.cdiv(m, 128), scale_n_pad)
    _quant_kernel[grid](
        a, q_u8, scale_u8,
        *a.stride(), *q_u8.stride(), *scale_u8.stride(),
        M=m, N=k, scaleN=scale_n_valid,
        scaleM_pad=scale_m_pad, scaleN_pad=scale_n_pad,
        BLOCK_SIZE=128, MXFP4_QUANT_BLOCK_SIZE=32,
        SCALING_MODE=0, SHUFFLE=True,
    )
    return q_u8.view(dtypes.fp4x2), scale_u8.view(dtypes.fp8_e8m0)


_quantize_a = _quantize_a_triton


def _get_out(device: torch.device, m: int, n: int):
    key = (device, m, n)
    out = _out_cache.get(key)
    if out is None:
        out = torch.empty((triton.cdiv(m, 32) * 32, n), dtype=dtypes.bf16, device=device)
        _out_cache[key] = out
    return out


def custom_kernel(data: input_t) -> output_t:
    a, _, _, b_shuffle, b_scale_sh = data
    m, k = a.shape
    n = b_shuffle.shape[0]

    a_q, a_scale_sh = _quantize_a(a)

    if gemm_a4w4_asm is None or get_GEMM_config is None:
        return aiter.gemm_a4w4(
            a_q, b_shuffle, a_scale_sh, b_scale_sh,
            dtype=dtypes.bf16, bpreshuffle=True,
        )

    out = _get_out(a_q.device, m, n)

    shape = (m, n, k)
    dispatch = _dispatch_cache.get(shape)
    if dispatch is None:
        config = get_GEMM_config(m, n, k)
        if config is not None:
            kn = config["kernelName"]
            sk = config.get("splitK", 0)
            sk = 0 if sk in (None, "") else int(sk)
            if kn and "_ZN" not in kn:
                dispatch = ("blockscale", kn, sk)
            else:
                dispatch = ("asm", kn, sk)
        elif k == 512 and m <= 8:
            dispatch = ("asm", _ASM_64X128, 0)
        elif m <= 64:
            dispatch = ("asm", _ASM_32X128, 0)
        else:
            dispatch = ("asm", "", 0)
        _dispatch_cache[shape] = dispatch

    kind, kernel_name, split_k = dispatch
    if kind == "blockscale":
        gemm_a4w4_blockscale(
            a_q.view(m, k // 2), b_shuffle, a_scale_sh, b_scale_sh,
            out, splitK=split_k,
        )
    else:
        if split_k:
            out.zero_()
        gemm_a4w4_asm(
            a_q.view(m, k // 2), 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]
scrolls · 134 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