Skip to content
KernelIndex
Search⌘K

submission 517208

migratesky · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:b4b2d51f2a267b9337605a2ec494823cbcac2b6216bd8661defd2e4f14000022
license declaredunknown
license concludedunknown
authorsmigratesky
imported2026-08-26

Techniques

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

fp4"""Shape-specialized submission for the AMD MXFP4 GEMM qualifier."""

Kernel source

amd_mxfp4_mm_submission.py127 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""Shape-specialized submission for the AMD MXFP4 GEMM qualifier."""

from collections import OrderedDict

import torch

from task import input_t, output_t

_AITER_STATE = None
_OUT_BUFFER_CACHE: OrderedDict[tuple[object, int, int], torch.Tensor] = OrderedDict()
_OUT_BUFFER_CACHE_MAX_ITEMS = 8

_ASM_32X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_ASM_64X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E"
_ASM_192X128 = "_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_192x128E"


def _init_aiter_state():
    global _AITER_STATE
    if _AITER_STATE is not None:
        return _AITER_STATE

    import aiter
    from aiter import QuantType, dtypes

    state = {
        "aiter": aiter,
        "dtypes": dtypes,
        "quant_func": aiter.get_triton_quant(QuantType.per_1x32),
    }

    try:
        from aiter.jit.utils.chip_info import get_cu_num
        from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm, get_GEMM_config

        # AITER's shipped tuned MXFP4 GEMM table only contains cu_num=256 rows.
        # Mirror those entries into the current CU bucket when needed so the
        # regular gemm_a4w4 wrapper can find tuned kernels instead of falling
        # back to the default selection path.
        get_GEMM_config(1, 512, 4096)
        gemm_dict = getattr(get_GEMM_config, "gemm_dict", None)
        if gemm_dict:
            current_cu = get_cu_num()
            if current_cu != 256:
                for (cu_num, m, n, k), config in list(gemm_dict.items()):
                    if cu_num != 256:
                        continue
                    gemm_dict.setdefault((current_cu, m, n, k), dict(config))
        state["gemm_a4w4_asm"] = gemm_a4w4_asm
    except Exception:
        pass

    _AITER_STATE = state
    return state


def _get_output_buffer(device, m, n):
    m_padded = (m + 31) // 32 * 32
    cache_key = (device, m_padded, n)
    out = _OUT_BUFFER_CACHE.get(cache_key)
    if out is not None:
        _OUT_BUFFER_CACHE.move_to_end(cache_key)
        return out

    out = torch.empty((m_padded, n), dtype=torch.bfloat16, device=device)
    _OUT_BUFFER_CACHE[cache_key] = out
    if len(_OUT_BUFFER_CACHE) > _OUT_BUFFER_CACHE_MAX_ITEMS:
        _OUT_BUFFER_CACHE.popitem(last=False)
    return out


def _pick_asm_kernel(m, n, k):
    if (n, k) == (2112, 7168):
        return _ASM_32X128, 0
    if (n, k) == (7168, 2048):
        return _ASM_32X128, 0
    if (n, k) == (3072, 1536):
        return _ASM_32X128, 0
    if k == 512 and n in (2880, 4096):
        return _ASM_64X128, 0
    return None



def custom_kernel(data: input_t) -> output_t:
    state = _init_aiter_state()

    a_src = data[0]
    b_shuffle = data[3]
    b_scale_sh = data[4]
    m, k = a_src.shape
    n = b_shuffle.shape[0]

    a_contiguous = a_src if a_src.is_contiguous() else a_src.contiguous()
    a_q, a_scale_sh = state["quant_func"](a_contiguous, shuffle=True)

    asm_config = _pick_asm_kernel(m, n, k)
    if asm_config is not None and "gemm_a4w4_asm" in state:
        kernel_name, log2_k_split = asm_config
        out = _get_output_buffer(a_q.device, m, n)
        state["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,
        )
        return out[:m]

    return state["aiter"].gemm_a4w4(
        a_q,
        b_shuffle,
        a_scale_sh,
        b_scale_sh,
        dtype=state["dtypes"].bf16,
        bpreshuffle=True,
    )
scrolls · 127 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