Skip to content
KernelIndex
Search⌘K

submission 754043

jkman2013 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-754043?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
#845 of 1143
2026-04-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:0c1021a2e8d941e72be03daf256260712db2939f290414cb6038258053277a5c
license declaredunknown
license concludedunknown
authorsjkman2013
imported2026-08-26

Techniques

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

fp4Optimized MXFP4 GEMM v3 for AMD MI355X.

Kernel source

submission_v3.py107 lines
"""
Optimized MXFP4 GEMM v3 for AMD MI355X.
- Direct asm kernel call (skip wrapper overhead)
- Pre-allocated output buffer + pad buffers in cache
- torch.inference_mode() to disable autograd overhead
- Module-level imports
"""
from task import input_t, output_t
import torch
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle

torch.set_grad_enabled(False)

_cache = {}
_asm_fn = None
_padm_fn = None
_initialized = False


def _init_asm():
    global _asm_fn, _padm_fn, _initialized
    _initialized = True
    try:
        _asm_fn = aiter.gemm_a4w4_asm
        _padm_fn = aiter.get_padded_m
    except AttributeError:
        _asm_fn = None


@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
    A, _, _, B_shuffle, B_scale_sh = data
    m, k = A.shape
    n = B_shuffle.shape[0]

    aq, asc = dynamic_mxfp4_quant(A)
    ash = e8m0_shuffle(asc)
    aq_v = aq.view(dtypes.fp4x2)
    ash_v = ash.view(dtypes.fp8_e8m0)

    if not _initialized:
        result = aiter.gemm_a4w4(
            aq_v, B_shuffle, ash_v, B_scale_sh,
            dtype=dtypes.bf16, bpreshuffle=True,
        )
        _init_asm()
        return result

    if _asm_fn is None:
        return aiter.gemm_a4w4(
            aq_v, B_shuffle, ash_v, B_scale_sh,
            dtype=dtypes.bf16, bpreshuffle=True,
        )

    key = (m, n, k)
    if key not in _cache:
        padded_m = _padm_fn(m, n, k, 1)
        kname = f"f4gemm_bf16_per1x32Fp4_BpreShuffle_{padded_m}x128"
        # Test if asm kernel exists for this shape
        try:
            out_buf = torch.empty((padded_m, n), dtype=torch.bfloat16, device="cuda")
            if m == padded_m:
                _asm_fn(aq_v, B_shuffle, ash_v, B_scale_sh, out_buf, kname, bpreshuffle=True)
            else:
                aq_pad = torch.zeros((padded_m, aq_v.shape[1]), dtype=aq_v.dtype, device="cuda")
                aq_pad[:m] = aq_v
                ash_pad = torch.zeros((padded_m, ash_v.shape[1]), dtype=ash_v.dtype, device="cuda")
                ash_pad[:m] = ash_v
                _asm_fn(aq_pad, B_shuffle, ash_pad, B_scale_sh, out_buf, kname, bpreshuffle=True)
            entry = {"out": out_buf, "kname": kname, "padded_m": padded_m, "use_asm": True}
            if m != padded_m:
                entry["aq_pad"] = aq_pad
                entry["ash_pad"] = ash_pad
            _cache[key] = entry
            return out_buf[:m] if m != padded_m else out_buf
        except RuntimeError:
            # Kernel not found for this shape, fall back to wrapper
            _cache[key] = {"use_asm": False}
            return aiter.gemm_a4w4(
                aq_v, B_shuffle, ash_v, B_scale_sh,
                dtype=dtypes.bf16, bpreshuffle=True,
            )

    c = _cache[key]
    if not c["use_asm"]:
        return aiter.gemm_a4w4(
            aq_v, B_shuffle, ash_v, B_scale_sh,
            dtype=dtypes.bf16, bpreshuffle=True,
        )

    out_buf = c["out"]
    padded_m = c["padded_m"]

    if m == padded_m:
        _asm_fn(aq_v, B_shuffle, ash_v, B_scale_sh, out_buf, c["kname"], bpreshuffle=True)
        return out_buf
    else:
        aq_pad = c["aq_pad"]
        ash_pad = c["ash_pad"]
        aq_pad[:m] = aq_v
        ash_pad[:m] = ash_v
        _asm_fn(aq_pad, B_shuffle, ash_pad, B_scale_sh, out_buf, c["kname"], bpreshuffle=True)
        return out_buf[:m]
scrolls · 107 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