Skip to content
KernelIndex
Search⌘K

submission 513839

jethreetwo · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

solution.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-513839?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.5µs
#442 of 1143
2026-03-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:14934c1eba45d5a4ad957f937a53fe668659d4c6a5c138152420bbd06a98506d
license declaredunknown
license concludedunknown
authorsjethreetwo
imported2026-08-26

Techniques

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

fp4Optimized MXFP4 matrix multiplication for AMD MI355X.

Kernel source

solution.py72 lines
"""
Optimized MXFP4 matrix multiplication for AMD MI355X.

Strategy:
- Keep the native per-1x32 quantization path
- Force a tuned asm kernel name directly for contest benchmark shapes
- Reuse the padded output buffer per shape
"""
import torch
from task import input_t, output_t
import aiter
from aiter import QuantType, dtypes

_QUANT_FUNC = aiter.get_triton_quant(QuantType.per_1x32)
_OUT_CACHE: dict = {}

_DEFAULT_KERNEL = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_KERNELS = {
    (4, 2880, 512): _DEFAULT_KERNEL,
    (16, 2112, 7168): _DEFAULT_KERNEL,
    (32, 4096, 512): _DEFAULT_KERNEL,
    (32, 2880, 512): _DEFAULT_KERNEL,
    (64, 7168, 2048): _DEFAULT_KERNEL,
    (256, 3072, 1536): _DEFAULT_KERNEL,
}


def _get_out(m: int, n: int, device: torch.device) -> torch.Tensor:
    key = (m, n, device.index)
    out = _OUT_CACHE.get(key)
    if out is None:
        padded_m = ((m + 31) // 32) * 32
        out = torch.empty((padded_m, 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
    A = A.contiguous()
    m, k = A.shape
    n = B_shuffle.shape[0]

    kernel_name = _KERNELS.get((m, n, k))
    A_q, A_scale_sh = _QUANT_FUNC(A, shuffle=True)

    if kernel_name is None:
        return aiter.gemm_a4w4(
            A_q,
            B_shuffle,
            A_scale_sh,
            B_scale_sh,
            dtype=dtypes.bf16,
            bpreshuffle=True,
        )

    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,
        None,
        1.0,
        0.0,
        True,
        log2_k_split=0,
    )
    return out[:m]
scrolls · 72 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