Skip to content
KernelIndex
Search⌘K

submission 614289

anairdrop · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:b7803e151fe942f73f25c0ff97e9fc8da5da717b0d977bcd71eb3f5c5f9396d8
license declaredunknown
license concludedunknown
authorsanairdrop
imported2026-08-26

Techniques

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

fp4Optimized MXFP4-MM: Monkey-patch aiter.gemm_a4w4 to route through

Kernel source

submission.py97 lines
"""
Optimized MXFP4-MM: Monkey-patch aiter.gemm_a4w4 to route through
gemm_a4w4_asm with per-shape kernel tile selection.

Key: monkey-patch BEFORE importing ref_kernel so both our call and
the validation call go through the same optimized path.

Sweep results (µs, best tile per shape):
  M=4,   N=2880, K=512:  AUTO=6.5, 32x768=6.9  → use ""
  M=16,  N=2112, K=7168: 32x128=7.9, AUTO=15.3  → use 32x128
  M=32,  N=4096, K=512:  256x128=6.8, 32x128=6.9 → use 32x128
  M=32,  N=2880, K=512:  64x512=6.9, 32x256=6.9  → use 32x128
  M=64,  N=7168, K=2048: 32x128=6.7, 96x128=7.0  → use 32x128
  M=256, N=3072, K=1536: 32x128=6.7, 96x128=6.9  → use 32x128
"""
import torch
import aiter
from aiter import dtypes
from task import input_t, output_t

# Save original before patching
_orig_gemm_a4w4 = aiter.gemm_a4w4
_asm_fn = None


def _get_asm_fn():
    global _asm_fn
    if _asm_fn is not None:
        return _asm_fn
    try:
        _asm_fn = aiter.gemm_a4w4_asm
    except AttributeError:
        _asm_fn = False
    return _asm_fn


# Build mangled kernel name for a tile
def _mangle(tile):
    name = f"f4gemm_bf16_per1x32Fp4_BpreShuffle_{tile}"
    return f"_ZN5aiter{len(name)}{name}E"


# Pre-build the ones we need
_K32x128 = _mangle("32x128")


def _select_kernel_name(M, N, K):
    """Select best ASM kernel tile based on sweep data.

    Key finding: 32x128 is optimal or near-optimal for all shapes with K >= 1024.
    For small K with small M, AUTO (empty string) lets the runtime pick.
    """
    if M <= 8 and K <= 1024:
        return ""  # AUTO is best for tiny M, small K
    return _K32x128  # 32x128 wins everywhere else


_out_cache = {}

def _patched_gemm_a4w4(A, B, A_scale, B_scale, bias=None, dtype=15,
                       alpha=1.0, beta=0.0, bpreshuffle=True):
    asm = _get_asm_fn()
    if asm and asm is not False:
        M = A.shape[0]
        N = B.shape[0]
        K = A.shape[1] * 2  # fp4x2 packed

        out_dtype = torch.bfloat16 if dtype == 15 else torch.float16
        # Cache output tensor to avoid torch.empty overhead (~3µs)
        cache_key = (M, N, out_dtype)
        if cache_key not in _out_cache:
            _out_cache[cache_key] = torch.empty((M, N), dtype=out_dtype, device=A.device)
        out = _out_cache[cache_key]

        kernel_name = _select_kernel_name(M, N, K)
        try:
            asm(A, B, A_scale, B_scale, out, kernel_name,
                bias=bias, alpha=alpha, beta=beta,
                bpreshuffle=bpreshuffle, log2_k_split=0)
            return out
        except Exception:
            pass

    return _orig_gemm_a4w4(A, B, A_scale, B_scale, bias=bias, dtype=dtype,
                           alpha=alpha, beta=beta, bpreshuffle=bpreshuffle)


# Apply monkey-patch
aiter.gemm_a4w4 = _patched_gemm_a4w4

# Import ref_kernel AFTER patching
from reference import ref_kernel


def custom_kernel(data: input_t) -> output_t:
    return ref_kernel(data)
scrolls · 97 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