Skip to content
KernelIndex
Search⌘K

submission 733679

Behzod12312121 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_2880_32x128_s3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-733679?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
22.4µs
#751 of 1143
2026-04-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:5329eedc53dc5d967d993b0f2f4dc32678314184c515e415672ac183970349e9
license declaredunknown
license concludedunknown
authorsBehzod12312121
imported2026-08-26

Techniques

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

split-kkernel as the working 2112 champion, with splitK=3 baked in.
tile-m = 1BLOCK_M: tl.constexpr = 1,

Kernel source

submission_2880_32x128_s3.py231 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
GEMM exact-shape route override for the 2880 family using the same 32x128 ASM
kernel as the working 2112 champion, with splitK=3 baked in.

Shapes overridden:
- (16, 2112, 7168) -> keep the known-good splitK=21 route
- (4, 2880, 512)   -> test 32x128 kernel with splitK from env
- (32, 2880, 512)  -> test 32x128 kernel with splitK from env

"""

import json
import os
import sys

import torch
import triton
import triton.language as tl

from task import input_t, output_t


_BENCHMARK_SHAPES = {
    (4, 2880, 512),
    (16, 2112, 7168),
    (32, 4096, 512),
    (32, 2880, 512),
    (64, 7168, 2048),
    (256, 3072, 1536),
}
_TARGET_2112 = (16, 2112, 7168)
_TARGET_2880_SHAPES = {
    (4, 2880, 512),
    (32, 2880, 512),
}
_KERNEL_32X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_SPLIT_K_2880 = 3
_TRACE_SEEN: set[tuple[int, int, int]] = set()
_PATCHED = False


def _emit_route(shape: tuple[int, int, int], config: dict | None, source: str) -> None:
    route = "asm"
    kernel_name = ""
    split_k = None
    if config is not None:
        kernel_name = str(config.get("kernelName", ""))
        split_k = config.get("splitK")
        if "_ZN" not in kernel_name:
            route = "blockscale"
    payload = {
        "shape": list(shape),
        "ck_config_found": config is not None,
        "kernelName": kernel_name,
        "splitK": split_k,
        "route": route,
        "source": source,
    }
    print(f"GEMM_TRACE {json.dumps(payload, sort_keys=True)}", file=sys.stderr)


def _install_route_patch() -> None:
    global _PATCHED
    if _PATCHED:
        return

    import aiter.ops.gemm_op_a4w4 as gemm_mod

    orig_get_config = gemm_mod.get_GEMM_config

    def wrapped_get_config(m: int, n: int, k: int):
        shape = (m, n, k)
        if shape == _TARGET_2112:
            config = {
                "kernelName": _KERNEL_32X128,
                "splitK": 21,
            }
            if shape not in _TRACE_SEEN:
                _emit_route(shape, config, source="manual_override_2112")
                _TRACE_SEEN.add(shape)
            return config
        if shape in _TARGET_2880_SHAPES:
            config = {
                "kernelName": _KERNEL_32X128,
                "splitK": _SPLIT_K_2880,
            }
            if shape not in _TRACE_SEEN:
                _emit_route(shape, config, source=f"manual_override_2880_32x128_s{_SPLIT_K_2880}")
                _TRACE_SEEN.add(shape)
            return config

        config = orig_get_config(m, n, k)
        if shape in _BENCHMARK_SHAPES and shape not in _TRACE_SEEN:
            _emit_route(shape, config, source="baseline")
            _TRACE_SEEN.add(shape)
        return config

    gemm_mod.get_GEMM_config = wrapped_get_config
    _PATCHED = True


@triton.jit
def _mxfp4_quant_kernel(
    x_ptr,
    out_ptr,
    scale_ptr,
    M,
    K,
    stride_xm,
    GROUP_SIZE: tl.constexpr = 32,
    BLOCK_M: tl.constexpr = 1,
):
    row = tl.program_id(0)
    group_id = tl.program_id(1)

    k_start = group_id * GROUP_SIZE
    half_offs = tl.arange(0, GROUP_SIZE // 2)

    k_even = k_start + half_offs * 2
    k_odd = k_start + half_offs * 2 + 1
    mask_e = k_even < K
    mask_o = k_odd < K

    x_even = tl.load(x_ptr + row * stride_xm + k_even, mask=mask_e, other=0.0).to(
        tl.float32
    )
    x_odd = tl.load(x_ptr + row * stride_xm + k_odd, mask=mask_o, other=0.0).to(
        tl.float32
    )

    abs_max = tl.maximum(
        tl.max(tl.abs(x_even), axis=0), tl.max(tl.abs(x_odd), axis=0)
    )

    abs_max = tl.maximum(abs_max, 1e-38).to(tl.float32)
    abs_max_int = abs_max.to(tl.int32, bitcast=True)
    abs_max_rounded = (
        (abs_max_int + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
    ).to(tl.float32, bitcast=True)
    scale_e8m0_unbiased = tl.floor(tl.math.log2(abs_max_rounded)).to(tl.int32) - 2
    scale_e8m0_unbiased = tl.minimum(
        tl.maximum(scale_e8m0_unbiased, -127), 127
    )
    e8m0_exp = (scale_e8m0_unbiased + 127).to(tl.uint8)
    quant_scale = tl.math.exp2(-scale_e8m0_unbiased.to(tl.float32))

    tl.store(scale_ptr + row * (K // GROUP_SIZE) + group_id, e8m0_exp)

    xs_e = x_even * quant_scale
    xs_e_uint = xs_e.to(tl.int32, bitcast=True).to(tl.uint32)
    s_e = xs_e_uint & 0x80000000
    xs_e_pos_uint = xs_e_uint ^ s_e
    xs_e_pos = xs_e_pos_uint.to(tl.float32, bitcast=True)
    sat_e = xs_e_pos >= 6.0
    den_e = xs_e_pos < 1.0
    mant_odd_e = (xs_e_pos_uint >> 22) & 1
    norm_e = (
        (xs_e_pos_uint.to(tl.int32) + (-1054867457)) + mant_odd_e.to(tl.int32)
    ) >> 22
    norm_e = norm_e.to(tl.uint8)
    den_val_e = (xs_e_pos + 4194304.0).to(tl.int32, bitcast=True) - 0x4A800000
    den_val_e = den_val_e.to(tl.uint8)
    q_e = tl.full(xs_e.shape, 7, dtype=tl.uint8)
    q_e = tl.where(~sat_e, norm_e, q_e)
    q_e = tl.where(den_e, den_val_e, q_e)
    sign_e_lp = (s_e >> 28).to(tl.uint8)
    lo = (q_e | sign_e_lp) & 0xF

    xs_o = x_odd * quant_scale
    xs_o_uint = xs_o.to(tl.int32, bitcast=True).to(tl.uint32)
    s_o = xs_o_uint & 0x80000000
    xs_o_pos_uint = xs_o_uint ^ s_o
    xs_o_pos = xs_o_pos_uint.to(tl.float32, bitcast=True)
    sat_o = xs_o_pos >= 6.0
    den_o = xs_o_pos < 1.0
    mant_odd_o = (xs_o_pos_uint >> 22) & 1
    norm_o = (
        (xs_o_pos_uint.to(tl.int32) + (-1054867457)) + mant_odd_o.to(tl.int32)
    ) >> 22
    norm_o = norm_o.to(tl.uint8)
    den_val_o = (xs_o_pos + 4194304.0).to(tl.int32, bitcast=True) - 0x4A800000
    den_val_o = den_val_o.to(tl.uint8)
    q_o = tl.full(xs_o.shape, 7, dtype=tl.uint8)
    q_o = tl.where(~sat_o, norm_o, q_o)
    q_o = tl.where(den_o, den_val_o, q_o)
    sign_o_lp = (s_o >> 28).to(tl.uint8)
    hi = ((q_o | sign_o_lp) & 0xF) << 4

    packed = lo | hi
    out_base = row * (K // 2) + k_start // 2
    tl.store(out_ptr + out_base + half_offs, packed.to(tl.uint8))


def _triton_mxfp4_quant(x: torch.Tensor):
    m, k = x.shape
    assert k % 32 == 0
    x = x.contiguous()

    out = torch.empty(m, k // 2, dtype=torch.uint8, device=x.device)
    scale = torch.empty(m, k // 32, dtype=torch.uint8, device=x.device)

    grid = (m, k // 32)
    _mxfp4_quant_kernel[grid](x, out, scale, m, k, x.stride(0))
    return out, scale


def custom_kernel(data: input_t) -> output_t:
    import aiter
    from aiter import dtypes
    from aiter.utility.fp4_utils import e8m0_shuffle

    _install_route_patch()

    a, _, _, b_shuffle, b_scale_sh = data
    a = a.contiguous()

    a_fp4, a_scale = _triton_mxfp4_quant(a)
    a_scale_sh = e8m0_shuffle(a_scale.view(torch.uint8))

    return aiter.gemm_a4w4(
        a_fp4.view(dtypes.fp4x2),
        b_shuffle,
        a_scale_sh.view(dtypes.fp8_e8m0),
        b_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )
scrolls · 231 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