Skip to content
KernelIndex
Search⌘K

submission 671219

windseeker.ws · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e9e9cf98a06b7f0bdb9fe5e9f57f2028825acdc3140a30f08b30401387f94d21
license declaredunknown
license concludedunknown
authorswindseeker.ws
imported2026-08-26

Techniques

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

fp4Aggressive MXFP4 MM submission:
split-kdef _run_asm_gemm(a_q, b_shuffle, a_scale_sh, b_scale_sh, kernel_name: str, split_k: int):

Kernel source

submission.py626 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
Aggressive MXFP4 MM submission:
- custom HIP kernel for A-side MXFP4 quantization
- direct production of packed fp4x2 A_q + shuffled E8M0 scales
- reuse aiter.gemm_a4w4 for the matrix multiply
"""

from __future__ import annotations

from dataclasses import dataclass
import os
import sys

from task import input_t, output_t


_OPTIMIZED_SHAPES = {
    (4, 2880, 512),
    (8, 2112, 7168),
    (16, 2112, 7168),
    (16, 3072, 1536),
    (32, 2880, 512),
    (32, 4096, 512),
    (64, 3072, 1536),
    (64, 7168, 2048),
    (256, 2880, 512),
    (256, 3072, 1536),
}

_EXTENSION_NAME = "mxfp4_mm_quant_v5"
_EXTENSION = None
_EXTENSION_FAILED = False
_REPORTED_MESSAGES: set[str] = set()
_PLAN_CACHE: dict[tuple[int, int, int], "_QuantPlan"] = {}
_BEST_GEMM_CACHE: dict[tuple[int, int, int, int], tuple[str, str | None, int]] = {}
_OUT_CACHE: dict[tuple[int, int, int], object] = {}

_ASM_KERNEL_32X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_ASM_KERNEL_64X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E"
_ASM_KERNEL_96X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_96x128E"
_ASM_KERNEL_128X128 = "_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_128x128E"
_ASM_KERNEL_192X128 = "_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_192x128E"
_ASM_KERNEL_224X128 = "_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_224x128E"
_ASM_KERNEL_256X256 = "_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_256x256E"

_ASM_CANDIDATES_BY_SHAPE: dict[tuple[int, int, int], tuple[str, ...]] = {
    (4, 2880, 512): (
        _ASM_KERNEL_192X128,
        _ASM_KERNEL_64X128,
        _ASM_KERNEL_32X128,
        _ASM_KERNEL_128X128,
        _ASM_KERNEL_256X256,
    ),
    (8, 2112, 7168): (
        _ASM_KERNEL_32X128,
        _ASM_KERNEL_224X128,
    ),
    (16, 2112, 7168): (
        _ASM_KERNEL_32X128,
        _ASM_KERNEL_224X128,
    ),
    (16, 3072, 1536): (
        _ASM_KERNEL_32X128,
        _ASM_KERNEL_64X128,
        _ASM_KERNEL_128X128,
    ),
    (32, 2880, 512): (
        _ASM_KERNEL_192X128,
        _ASM_KERNEL_64X128,
        _ASM_KERNEL_32X128,
        _ASM_KERNEL_128X128,
        _ASM_KERNEL_256X256,
    ),
    (32, 4096, 512): (
        _ASM_KERNEL_192X128,
        _ASM_KERNEL_256X256,
        _ASM_KERNEL_64X128,
        _ASM_KERNEL_32X128,
        _ASM_KERNEL_128X128,
    ),
    (64, 3072, 1536): (
        _ASM_KERNEL_32X128,
        _ASM_KERNEL_64X128,
        _ASM_KERNEL_128X128,
    ),
    (64, 7168, 2048): (
        _ASM_KERNEL_32X128,
        _ASM_KERNEL_64X128,
        _ASM_KERNEL_96X128,
        _ASM_KERNEL_128X128,
    ),
    (256, 2880, 512): (
        _ASM_KERNEL_256X256,
        _ASM_KERNEL_192X128,
        _ASM_KERNEL_64X128,
        _ASM_KERNEL_32X128,
    ),
    (256, 3072, 1536): (
        _ASM_KERNEL_32X128,
        _ASM_KERNEL_64X128,
        _ASM_KERNEL_128X128,
    ),
}


@dataclass
class _QuantPlan:
    a_q_u8: object
    a_q_fp4: object
    a_scale_sh_u8: object
    a_scale_sh_e8m0: object


_CPP_WRAPPER = """
void quantize_a_mxfp4(torch::Tensor a,
                      torch::Tensor a_q,
                      torch::Tensor a_scale_sh);
"""

_HIP_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>

#include <cstdint>
#include <stdexcept>

namespace {

__device__ __forceinline__ uint8_t float_to_fp4_e2m1_bits(float x) {
    uint32_t bits = __float_as_uint(x);
    uint32_t sign = bits & 0x80000000u;
    bits ^= sign;
    float x_abs = __uint_as_float(bits);

    uint8_t code = 0;
    if (x_abs >= 6.0f) {
        code = 0x7u;
    } else {
        if (x_abs < 1.0f) {
            constexpr int denorm_exp = ((127 - 1) + (23 - 1) + 1);
            constexpr uint32_t denorm_mask_int = static_cast<uint32_t>(denorm_exp) << 23;
            float denorm_mask_f32 = __uint_as_float(denorm_mask_int);
            uint32_t denorm_bits = __float_as_uint(x_abs + denorm_mask_f32) - denorm_mask_int;
            code = static_cast<uint8_t>(denorm_bits);
        } else {
            uint32_t normal_bits = bits;
            uint32_t mant_odd = (normal_bits >> (23 - 1)) & 1u;
            constexpr int32_t round_bias =
                ((1 - 127) << 23) + (1 << (23 - 2)) - 1;
            int32_t rounded = static_cast<int32_t>(normal_bits);
            rounded += round_bias;
            rounded += static_cast<int32_t>(mant_odd);
            code = static_cast<uint8_t>(static_cast<uint32_t>(rounded) >> (23 - 1));
        }
    }

    uint8_t sign_lp = static_cast<uint8_t>(sign >> 28);
    return static_cast<uint8_t>(code | sign_lp);
}

__global__ __launch_bounds__(32) void quantize_a_mxfp4_kernel(
    const __hip_bfloat16* __restrict__ a,
    uint8_t* __restrict__ a_q,
    uint8_t* __restrict__ a_scale_sh,
    int m,
    int k,
    int scale_cols
) {
    const int row = static_cast<int>(blockIdx.x);
    const int block64 = static_cast<int>(blockIdx.y);
    const int lane = static_cast<int>(threadIdx.x);
    const int subgroup = lane >> 4;
    const int sublane = lane & 15;

    if (row >= m) {
        return;
    }

    const int k_base = block64 * 64 + subgroup * 32 + sublane * 2;
    const float x0 = static_cast<float>(a[row * k + k_base]);
    const float x1 = static_cast<float>(a[row * k + k_base + 1]);
    float local_abs = fmaxf(fabsf(x0), fabsf(x1));
    for (int offset = 8; offset > 0; offset >>= 1) {
        local_abs = fmaxf(local_abs, __shfl_down(local_abs, offset, 16));
    }

    const float amax = __shfl(local_abs, 0, 16);
    uint8_t scale_byte = 0u;
    float quant_scale = 0.0f;
    if (sublane == 0) {
        if (amax > 0.0f) {
            uint32_t amax_bits = __float_as_uint(amax);
            amax_bits = (amax_bits + 0x200000u) & 0xFF800000u;
            const int scale_unbiased = static_cast<int>((amax_bits >> 23) & 0xFFu) - 127 - 2;
            scale_byte = static_cast<uint8_t>(scale_unbiased + 127);
            quant_scale = ldexpf(1.0f, -scale_unbiased);
        }
    }
    scale_byte = static_cast<uint8_t>(__shfl(static_cast<int>(scale_byte), 0, 16));
    quant_scale = __shfl(quant_scale, 0, 16);

    uint8_t code0 = 0u;
    uint8_t code1 = 0u;
    if (scale_byte != 0u) {
        const float scale = quant_scale;
        code0 = float_to_fp4_e2m1_bits(x0 * scale);
        code1 = float_to_fp4_e2m1_bits(x1 * scale);
    }
    const int q_col = block64 * 32 + subgroup * 16 + sublane;
    a_q[row * (k >> 1) + q_col] = static_cast<uint8_t>(code0 | (code1 << 4));

    if (sublane == 0) {
        const int raw_block = block64 * 2 + subgroup;
        const int a_tile = row >> 5;
        const int b = (row >> 4) & 1;
        const int c = row & 15;
        const int d = raw_block >> 3;
        const int e = (raw_block >> 2) & 1;
        const int f = raw_block & 3;
        const int d_tiles = scale_cols >> 3;
        const int shuffled_idx =
            (((((a_tile * d_tiles + d) * 4 + f) * 16 + c) * 2 + e) * 2 + b);
        a_scale_sh[shuffled_idx] = scale_byte;
    }
}

__global__ void quantize_a_mxfp4_row_kernel(
    const __hip_bfloat16* __restrict__ a,
    uint8_t* __restrict__ a_q,
    uint8_t* __restrict__ a_scale_sh,
    int m,
    int k,
    int scale_cols
) {
    const int row = static_cast<int>(blockIdx.x);
    const int lane = static_cast<int>(threadIdx.x);
    const int subgroup = lane >> 4;
    const int sublane = lane & 15;

    if (row >= m) {
        return;
    }

    const int q_cols = k >> 1;
    if (lane >= q_cols) {
        return;
    }

    const int k_base = lane * 2;
    const float x0 = static_cast<float>(a[row * k + k_base]);
    const float x1 = static_cast<float>(a[row * k + k_base + 1]);
    float local_abs = fmaxf(fabsf(x0), fabsf(x1));
    for (int offset = 8; offset > 0; offset >>= 1) {
        local_abs = fmaxf(local_abs, __shfl_down(local_abs, offset, 16));
    }

    const float amax = __shfl(local_abs, 0, 16);
    uint8_t scale_byte = 0u;
    float quant_scale = 0.0f;
    if (sublane == 0) {
        if (amax > 0.0f) {
            uint32_t amax_bits = __float_as_uint(amax);
            amax_bits = (amax_bits + 0x200000u) & 0xFF800000u;
            const int scale_unbiased = static_cast<int>((amax_bits >> 23) & 0xFFu) - 127 - 2;
            scale_byte = static_cast<uint8_t>(scale_unbiased + 127);
            quant_scale = ldexpf(1.0f, -scale_unbiased);
        }
    }
    scale_byte = static_cast<uint8_t>(__shfl(static_cast<int>(scale_byte), 0, 16));
    quant_scale = __shfl(quant_scale, 0, 16);

    uint8_t code0 = 0u;
    uint8_t code1 = 0u;
    if (scale_byte != 0u) {
        code0 = float_to_fp4_e2m1_bits(x0 * quant_scale);
        code1 = float_to_fp4_e2m1_bits(x1 * quant_scale);
    }
    a_q[row * q_cols + lane] = static_cast<uint8_t>(code0 | (code1 << 4));

    if (sublane == 0) {
        const int raw_block = subgroup;
        const int a_tile = row >> 5;
        const int b = (row >> 4) & 1;
        const int c = row & 15;
        const int d = raw_block >> 3;
        const int e = (raw_block >> 2) & 1;
        const int f = raw_block & 3;
        const int d_tiles = scale_cols >> 3;
        const int shuffled_idx =
            (((((a_tile * d_tiles + d) * 4 + f) * 16 + c) * 2 + e) * 2 + b);
        a_scale_sh[shuffled_idx] = scale_byte;
    }
}

}  // namespace

void quantize_a_mxfp4(torch::Tensor a,
                      torch::Tensor a_q,
                      torch::Tensor a_scale_sh) {
    TORCH_CHECK(a.is_cuda(), "A must be CUDA/HIP tensor");
    TORCH_CHECK(a_q.is_cuda(), "A_q must be CUDA/HIP tensor");
    TORCH_CHECK(a_scale_sh.is_cuda(), "A_scale_sh must be CUDA/HIP tensor");
    TORCH_CHECK(a.dim() == 2, "A must be 2D");
    TORCH_CHECK(a.scalar_type() == at::kBFloat16, "A must be bf16");
    TORCH_CHECK(a_q.scalar_type() == at::kByte, "A_q must be uint8");
    TORCH_CHECK(a_scale_sh.scalar_type() == at::kByte, "A_scale_sh must be uint8");

    const int64_t m = a.size(0);
    const int64_t k = a.size(1);
    TORCH_CHECK((k % 64) == 0, "k must be divisible by 64");
    TORCH_CHECK(a_q.size(0) == m, "A_q rows mismatch");
    TORCH_CHECK(a_q.size(1) == (k / 2), "A_q cols mismatch");
    TORCH_CHECK(a_scale_sh.size(0) == 256, "A_scale_sh rows must be 256");
    TORCH_CHECK(a_scale_sh.size(1) == (k / 32), "A_scale_sh cols mismatch");

    const int scale_cols = static_cast<int>(k / 32);
    if (k <= 2048) {
        dim3 grid(static_cast<unsigned int>(m));
        dim3 block(static_cast<unsigned int>(k / 2));
        quantize_a_mxfp4_row_kernel<<<grid, block, 0, 0>>>(
            reinterpret_cast<const __hip_bfloat16*>(a.data_ptr()),
            reinterpret_cast<uint8_t*>(a_q.data_ptr()),
            reinterpret_cast<uint8_t*>(a_scale_sh.data_ptr()),
            static_cast<int>(m),
            static_cast<int>(k),
            scale_cols
        );
    } else {
        dim3 grid(static_cast<unsigned int>(m), static_cast<unsigned int>(k / 64));
        dim3 block(32);
        quantize_a_mxfp4_kernel<<<grid, block, 0, 0>>>(
            reinterpret_cast<const __hip_bfloat16*>(a.data_ptr()),
            reinterpret_cast<uint8_t*>(a_q.data_ptr()),
            reinterpret_cast<uint8_t*>(a_scale_sh.data_ptr()),
            static_cast<int>(m),
            static_cast<int>(k),
            scale_cols
        );
    }

    hipError_t err = hipGetLastError();
    if (err != hipSuccess) {
        throw std::runtime_error(hipGetErrorString(err));
    }
}
"""


def _report_once(message: str):
    if message in _REPORTED_MESSAGES:
        return
    _REPORTED_MESSAGES.add(message)
    print(message, file=sys.stderr, flush=True)


def _reference_impl(data: input_t):
    import aiter
    from aiter import dtypes
    from aiter.ops.triton.quant import dynamic_mxfp4_quant
    from aiter.utility.fp4_utils import e8m0_shuffle

    a, _, _, b_shuffle, b_scale_sh = data
    a = a.contiguous()
    a_q, a_scale = dynamic_mxfp4_quant(a)
    a_scale_sh = e8m0_shuffle(a_scale)
    return aiter.gemm_a4w4(
        a_q.view(dtypes.fp4x2),
        b_shuffle,
        a_scale_sh.view(dtypes.fp8_e8m0),
        b_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )


def _get_extension():
    global _EXTENSION, _EXTENSION_FAILED

    if _EXTENSION is not None:
        return _EXTENSION
    if _EXTENSION_FAILED:
        return None

    import torch
    from torch.utils.cpp_extension import load_inline

    os.environ.setdefault("CXX", "clang++")
    os.environ.setdefault("TORCH_EXTENSIONS_DIR", "/tmp/torch_extensions")

    try:
        _EXTENSION = load_inline(
            name=_EXTENSION_NAME,
            cpp_sources=[_CPP_WRAPPER],
            cuda_sources=[_HIP_SRC],
            functions=["quantize_a_mxfp4"],
            verbose=False,
            extra_cuda_cflags=[
                "-O3",
                "-std=c++20",
                "--offload-arch=gfx950",
                "--offload-arch=gfx942",
            ],
        )
        return _EXTENSION
    except Exception as exc:
        _report_once(f"[mxfp4-mm] custom quant extension unavailable: {exc}")
        _EXTENSION_FAILED = True
        return None


def _get_gemm_backend():
    try:
        from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
        return gemm_a4w4_asm
    except Exception as exc:
        _report_once(f"[mxfp4-mm] low-level gemm unavailable: {exc}")
        return None


def _lookup_gemm_config(m: int, n: int, k: int):
    try:
        from aiter.ops.gemm_op_a4w4 import get_GEMM_config
        return get_GEMM_config(m, n, k)
    except Exception as exc:
        _report_once(f"[mxfp4-mm] gemm config lookup unavailable: {exc}")
        return None


def _get_out_buffer(device, m: int, n: int):
    import torch

    device_index = -1 if device.index is None else device.index
    key = (device_index, m, n)
    out = _OUT_CACHE.get(key)
    if out is None or out.device != device:
        padded_m = (m + 31) // 32 * 32
        out = torch.empty((padded_m, n), dtype=torch.bfloat16, device=device)
        _OUT_CACHE[key] = out
    return out


def _get_quant_plan(a):
    import torch
    from aiter import dtypes

    m, k = a.shape
    device_index = -1 if a.device.index is None else a.device.index
    key = (device_index, m, k)
    plan = _PLAN_CACHE.get(key)
    if plan is not None:
        return plan

    scale_cols = k // 32
    a_q_u8 = torch.empty((m, k // 2), dtype=torch.uint8, device=a.device)
    a_scale_sh_u8 = torch.empty((256, scale_cols), dtype=torch.uint8, device=a.device)
    plan = _QuantPlan(
        a_q_u8=a_q_u8,
        a_q_fp4=a_q_u8.view(dtypes.fp4x2),
        a_scale_sh_u8=a_scale_sh_u8,
        a_scale_sh_e8m0=a_scale_sh_u8.view(dtypes.fp8_e8m0),
    )
    _PLAN_CACHE[key] = plan
    return plan


def _run_standard_gemm(a_q, b_shuffle, a_scale_sh, b_scale_sh):
    import aiter
    from aiter import dtypes

    return aiter.gemm_a4w4(
        a_q,
        b_shuffle,
        a_scale_sh,
        b_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )


def _run_asm_gemm(a_q, b_shuffle, a_scale_sh, b_scale_sh, kernel_name: str, split_k: int):
    gemm_a4w4_asm = _get_gemm_backend()
    if gemm_a4w4_asm is None:
        raise RuntimeError("gemm_a4w4_asm unavailable")

    m = a_q.shape[0]
    n = b_shuffle.shape[0]
    out = _get_out_buffer(a_q.device, m, n)
    gemm_a4w4_asm(
        a_q,
        b_shuffle,
        a_scale_sh,
        b_scale_sh,
        out,
        kernelName=kernel_name,
        bias=None,
        alpha=1.0,
        beta=0.0,
        bpreshuffle=True,
        log2_k_split=split_k,
    )
    return out[:m]


def _select_best_gemm(shape_key, a_q, b_shuffle, a_scale_sh, b_scale_sh):
    import torch

    cached = _BEST_GEMM_CACHE.get(shape_key)
    if cached is not None:
        return cached

    _, m, n, k = shape_key
    config = _lookup_gemm_config(m, n, k)
    if config is not None:
        best_choice = ("standard", None, 0)
        _BEST_GEMM_CACHE[shape_key] = best_choice
        _report_once(f"[mxfp4-mm] gemm choice for {(m, n, k)} -> tuned standard")
        return best_choice

    baseline = _run_standard_gemm(a_q, b_shuffle, a_scale_sh, b_scale_sh)
    candidates = [("standard", None, 0, lambda: _run_standard_gemm(a_q, b_shuffle, a_scale_sh, b_scale_sh))]

    gemm_a4w4_asm = _get_gemm_backend()
    if gemm_a4w4_asm is not None:
        for kernel_name in _ASM_CANDIDATES_BY_SHAPE.get(shape_key[1:], (_ASM_KERNEL_32X128,)):
            candidates.append(
                (
                    "asm",
                    kernel_name,
                    0,
                    lambda kernel_name=kernel_name: _run_asm_gemm(
                        a_q,
                        b_shuffle,
                        a_scale_sh,
                        b_scale_sh,
                        kernel_name,
                        0,
                    ),
                )
            )

    best_choice = ("standard", None, 0)
    best_time = float("inf")
    for mode, kernel_name, split_k, fn in candidates:
        try:
            candidate_out = fn()
            if not torch.allclose(candidate_out, baseline, rtol=1e-2, atol=1e-2):
                continue

            elapsed_us = float("inf")
            for _ in range(3):
                torch.cuda.synchronize()
                start = torch.cuda.Event(enable_timing=True)
                end = torch.cuda.Event(enable_timing=True)
                start.record()
                fn()
                end.record()
                torch.cuda.synchronize()
                elapsed_us = min(elapsed_us, start.elapsed_time(end))
            if elapsed_us < best_time:
                best_time = elapsed_us
                best_choice = (mode, kernel_name, split_k)
        except Exception as exc:
            _report_once(f"[mxfp4-mm] gemm candidate skipped ({mode}, {kernel_name}): {exc}")

    _BEST_GEMM_CACHE[shape_key] = best_choice
    _report_once(f"[mxfp4-mm] gemm choice for {(m, n, k)} -> {best_choice[0]} {best_choice[1] or 'aiter'}")
    return best_choice


def _optimized_impl(data: input_t):
    module = _get_extension()
    if module is None:
        return None

    a, _, _, b_shuffle, b_scale_sh = data
    a = a.contiguous()
    m, k = a.shape
    n = b_shuffle.shape[0]
    if (m, n, k) not in _OPTIMIZED_SHAPES:
        return None

    plan = _get_quant_plan(a)
    try:
        module.quantize_a_mxfp4(a, plan.a_q_u8, plan.a_scale_sh_u8)
    except Exception as exc:
        _report_once(f"[mxfp4-mm] custom quant kernel fallback: {exc}")
        return None

    shape_key = (-1 if a.device.index is None else a.device.index, m, n, k)
    mode, kernel_name, split_k = _select_best_gemm(
        shape_key,
        plan.a_q_fp4,
        b_shuffle,
        plan.a_scale_sh_e8m0,
        b_scale_sh,
    )
    if mode == "asm" and kernel_name is not None:
        try:
            return _run_asm_gemm(
                plan.a_q_fp4,
                b_shuffle,
                plan.a_scale_sh_e8m0,
                b_scale_sh,
                kernel_name,
                split_k,
            )
        except Exception as exc:
            _report_once(f"[mxfp4-mm] chosen asm gemm fallback: {exc}")

    return _run_standard_gemm(
        plan.a_q_fp4,
        b_shuffle,
        plan.a_scale_sh_e8m0,
        b_scale_sh,
    )


def custom_kernel(data: input_t) -> output_t:
    out = _optimized_impl(data)
    if out is not None:
        return out
    return _reference_impl(data)
scrolls · 626 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