Skip to content
KernelIndex
Search⌘K

submission 721832

nanbeilvdougao · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_20260404_16.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-721832?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.0µs
#719 of 1143
2026-04-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:40ac52dcb95f67a2eaa9e08e95a8a763ff89f8c095bc733138d5f67be47f6d27
license declaredunknown
license concludedunknown
authorsnanbeilvdougao
imported2026-08-26

Techniques

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

split-kkname, splitk = override

Kernel source

submission_20260404_16.py64 lines
from __future__ import annotations

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
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4

from task import input_t, output_t

_K32 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_FORCE = {
    (4, 2880, 512): (_K32, 0),
    (16, 2112, 7168): (_K32, 0),
    (32, 4096, 512): (_K32, 0),
    (32, 2880, 512): (_K32, 0),
}


def _quantize_activation(a: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    a_bf16 = a.contiguous().to(dtype=torch.bfloat16)
    a_q, a_scale = dynamic_mxfp4_quant(a_bf16)
    a_scale_sh = e8m0_shuffle(a_scale)
    return a_q.view(dtypes.fp4x2), a_scale_sh.view(dtypes.fp8_e8m0)


def _best_known_wrapper_hybrid(a: torch.Tensor, b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor) -> torch.Tensor:
    a_q, a_scale_sh = _quantize_activation(a)
    m = a.shape[0]
    n = b_shuffle.shape[0]
    k = a.shape[1]
    override = _FORCE.get((m, n, k))
    if override is not None:
        kname, splitk = override
        out = torch.empty(((m + 31) // 32) * 32, n, dtype=torch.bfloat16, device=a.device)
        aiter.gemm_a4w4_asm(a_q.view(m, k // 2), b_shuffle, a_scale_sh, b_scale_sh, out, kname, bpreshuffle=True, log2_k_split=splitk)
        return out[:m, :n]
    out = aiter.gemm_a4w4(a_q, b_shuffle, a_scale_sh, b_scale_sh, dtype=dtypes.bf16, bpreshuffle=True)
    return out[:m, :n]


def _try_dense_requant_plain_all(a: torch.Tensor, b_dense: torch.Tensor) -> torch.Tensor | None:
    a_bf16 = a.contiguous().to(dtype=torch.bfloat16)
    m, k = a_bf16.shape
    n = b_dense.shape[0]
    if k % 64 != 0:
        return None
    b_dense = b_dense.contiguous().to(dtype=torch.bfloat16)
    b_q, b_scale = dynamic_mxfp4_quant(b_dense)
    out = gemm_a16wfp4(a_bf16, b_q.view(n, k // 2), b_scale.view(n, k // 32), dtype=torch.bfloat16)
    return out[:m, :n]


def custom_kernel(data: input_t) -> output_t:
    a, b_dense, _b_q, b_shuffle, b_scale_sh = data
    try:
        out = _try_dense_requant_plain_all(a, b_dense)
        if out is not None:
            return out
    except Exception:
        pass
    return _best_known_wrapper_hybrid(a, b_shuffle, b_scale_sh)
scrolls · 64 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