Skip to content
KernelIndex
Search⌘K

submission 720595

yszheda · 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.

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f190666b85611495fb5d95e278c33c388f8a7607e21cbdb7be9b2d6a96a5eb6f
license declaredunknown
license concludedunknown
authorsyszheda
imported2026-08-26

Techniques

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

fp4a_dtype="fp4",

Kernel source

submission.py127 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
import sys
import os

try:
    import flydsl

    _flydsl_root = os.path.dirname(os.path.dirname(flydsl.__file__))
    if _flydsl_root not in sys.path:
        sys.path.insert(0, _flydsl_root)
except Exception:
    pass


def _patch_triton_fp4():
    try:
        import triton
        import triton.language as tl

        if hasattr(tl, "core") and hasattr(tl.core, "dtype"):
            if not hasattr(tl.core.dtype, "SUPPORTED_TENSOR_DTYPES"):
                tl.core.dtype.SUPPORTED_TENSOR_DTYPES = set()
            tl.core.dtype.SUPPORTED_TENSOR_DTYPES.add("float4_e2m1fn_x2")
        try:
            import triton._utils as _tu

            if hasattr(_tu, "type_canonicalisation_dict"):
                _tu.type_canonicalisation_dict["float4_e2m1fn_x2"] = (
                    "*kfloat4_e2m1fn_x2"
                )
        except Exception:
            pass
    except Exception:
        pass


_patch_triton_fp4()

import torch
from task import input_t, output_t
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle


_flydsl_available = False
try:
    from kernels.preshuffle_gemm import compile_preshuffle_gemm_w4

    _flydsl_available = True
except Exception:
    pass

_TILE_CONFIGS = [
    (8, 32, 128, 256, False, None),
    (16, 32, 128, 256, False, None),
    (32, 32, 128, 256, False, None),
    (64, 64, 256, 256, False, None),
    (128, 64, 256, 256, False, None),
    (256, 64, 256, 256, False, None),
    (512, 64, 256, 256, False, None),
    (float("inf"), 64, 256, 256, False, None),
]


def _get_tile_config(M, N, K):
    for m_max, tm, tn, tk, async_copy, wpe in _TILE_CONFIGS:
        if M <= m_max:
            return tm, tn, tk, async_copy, wpe
    return 64, 256, 256, False, None


_kernel_cache = {}


def _get_cached_kernel(M, N, K):
    cache_key = (M, N, K)
    if cache_key in _kernel_cache:
        return _kernel_cache[cache_key]

    tile_m, tile_n, tile_k, use_async_copy, waves_per_eu = _get_tile_config(M, N, K)

    launch_fn = compile_preshuffle_gemm_w4(
        M=M,
        N=N,
        K=K,
        tile_m=tile_m,
        tile_n=tile_n,
        tile_k=tile_k,
        a_dtype="fp4",
        b_dtype="fp4",
        out_dtype="bf16",
        lds_stage=2,
        use_cshuffle_epilog=False,
        waves_per_eu=waves_per_eu,
        use_async_copy=use_async_copy,
        dsrd_preload=2,
        dvmem_preload=2,
    )

    _kernel_cache[cache_key] = launch_fn
    return launch_fn


def _quant_mxfp4(x, shuffle=True):
    x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
    if shuffle:
        bs_e8m0 = e8m0_shuffle(bs_e8m0)
    return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)


def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    m, k = A.shape
    n, _ = B.shape

    A_q, A_scale_sh = _quant_mxfp4(A, shuffle=True)
    return torch.ops.aiter.gemm_a4w4(
        A_q,
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        dtype=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