Skip to content
KernelIndex
Search⌘K

submission 745884

Xerous Wazler · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:0f20ddce76c078d6672c235cff6e9629ec77fcdecec9985cb712b19f6cac5bdc
license declaredunknown
license concludedunknown
authorsXerous Wazler
imported2026-08-26

Techniques

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

fp4FP4 quant + FP4 GEMM optimized for MI355X:

Kernel source

submission_v1.py188 lines
"""
FP4 quant + FP4 GEMM optimized for MI355X:
  bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.

Replaces dynamic_mxfp4_quant + e8m0_shuffle (2 kernel launches, ~10-12µs)
with a single fused HIP kernel (~2-3µs target).
"""
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

from task import input_t, output_t

# ── Build HIP fused quant+shuffle kernel ────────────────────────────
_hip_mod = None

try:
    from torch.utils.cpp_extension import load_inline

   
    _hip_kernel_code = r"""

#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_warp_functions.h>

#define WAVE_SIZE 32
#define GROUP_SIZE 32
#define GROUPS_PER_BLOCK 2

// FP4 E2M1 quantization thresholds (scaled)
__device__ __forceinline__ unsigned int float_to_fp4_code(float x) {
    if (x < 0.25f) return 0;
    if (x < 0.75f) return 1;
    if (x < 1.25f) return 2;
    if (x < 1.75f) return 3;
    if (x < 2.5f)  return 4;
    if (x < 3.5f)  return 5;
    if (x < 5.0f)  return 6;
    return 7;
}

__global__ __launch_bounds__(64, 4)
void fused_mxfp4_quant_e8m0shuffle_latency(
    const unsigned short* __restrict__ input_bf16, // [M, K]
    unsigned char*        __restrict__ out_fp4,    // [M, K/2]
    unsigned char*        __restrict__ out_scale,  // [padM, K/32]
    int M, int K, int n_groups
) {
    int tid  = threadIdx.x;
    int wave = tid >> 5;            // 0 or 1
    int lane = tid & 31;            // 0..31

    int row = blockIdx.y;
    int grp = blockIdx.x * GROUPS_PER_BLOCK + wave;

    if (row >= M || grp >= n_groups)
        return;

    int base_k = grp * GROUP_SIZE + lane;
    int idx    = row * K + base_k;

    // ---- BF16 -> F32 ----
    unsigned short bf = input_bf16[idx];
    float fval = __uint_as_float(((unsigned int)bf) << 16);
    float aval = fabsf(fval);

    // ---- Wavefront max reduction ----
    float amax = __builtin_amdgcn_wavefront_max(aval);

    // ---- Lane 0 computes E8M0 + inv_scale ----
    unsigned int e8m0 = 0;
    unsigned int inv_scale_bits = 0;

    if (lane == 0) {
        if (amax != 0.0f) {
            unsigned int bits = __float_as_uint(amax);
            e8m0 = (bits >> 23) & 0xFF;
            unsigned int inv_exp = 254u - e8m0;
            if (inv_exp > 0 && inv_exp < 255)
                inv_scale_bits = inv_exp << 23;
        }
    }

    // Broadcast scale
    e8m0 = __builtin_amdgcn_readfirstlane(e8m0);
    inv_scale_bits = __builtin_amdgcn_readfirstlane(inv_scale_bits);
    float inv_scale = __uint_as_float(inv_scale_bits);

    // ---- Quantize ----
    float q = fabsf(fval) * inv_scale;
    if (q > 6.0f) q = 6.0f;

    unsigned int fp4 = float_to_fp4_code(q);
    if (fval < 0.0f) fp4 |= 8u;

    // ---- FP4 pack using lane ops ----
    if ((lane & 1) == 0) {
        unsigned int hi = __builtin_amdgcn_ds_bpermute((lane + 1) * 4, fp4);
        unsigned char packed = (fp4 & 0xF) | ((hi & 0xF) << 4);

        int out_k = grp * 16 + (lane >> 1);
        out_fp4[row * (K >> 1) + out_k] = packed;
    }

    // ---- Write shuffled E8M0 scale (lane 0 only) ----
    if (lane == 0) {
        int rem = grp & 3;
        int base = grp & ~3;
        int shuffled = base | ((rem == 1) ? 2 : (rem == 2) ? 1 : rem);
        out_scale[row * n_groups + shuffled] = (unsigned char)e8m0;
    }
}

// ---- C++ binding ----
std::vector<at::Tensor> fused_quant_fn(at::Tensor A) {
    TORCH_CHECK(A.is_cuda(), "CUDA tensor expected");
    TORCH_CHECK(A.scalar_type() == at::kBFloat16, "bf16 required");
    TORCH_CHECK(A.dim() == 2, "A must be 2D");

    int M = A.size(0);
    int K = A.size(1);
    TORCH_CHECK(K % 32 == 0, "K must be divisible by 32");

    int n_groups = K / 32;
    int padM = ((M + 31) / 32) * 32;

    auto opts = A.options().dtype(at::kByte);
    auto fp4_out = at::empty({M, K / 2}, opts);
    auto scale_out = at::zeros({padM, n_groups}, opts);

    dim3 block(64);
    dim3 grid((n_groups + 1) / 2, M);

    hipLaunchKernelGGL(
        fused_mxfp4_quant_e8m0shuffle_latency,
        grid, block, 0, 0,
        reinterpret_cast<const unsigned short*>(A.data_ptr<at::BFloat16>()),
        fp4_out.data_ptr<unsigned char>(),
        scale_out.data_ptr<unsigned char>(),
        M, K, n_groups
    );

    return {fp4_out, scale_out};
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
    m.def("fused_quant_fn", &fused_quant_fn, "Latency-optimized MXFP4 quant");
}
"""

    _hip_mod = load_inline(
        name="fused_mxfp4_hip_v4",
        cpp_sources="",
        cuda_sources=_hip_kernel_code,
        functions=["fused_quant_fn"],
        extra_cuda_cflags=["-O3"],
        verbose=False,
    )
    print("[submission] HIP fused quant+shuffle kernel compiled successfully")
except Exception as e:
    import traceback
    traceback.print_exc()
    print(f"[submission] HIP compilation failed: {e}")
    _hip_mod = None


def custom_kernel(data: input_t) -> output_t:
    A, _B, _B_q, B_shuffle, B_scale_sh = data

    if _hip_mod is not None:
        # ── HIP path: single fused kernel ──
        results = _hip_mod.fused_quant_fn(A)
        A_q = results[0].view(dtypes.fp4x2)
        A_scale = results[1].view(dtypes.fp8_e8m0)
    else:
        # ── Triton fallback ──
        A_q, A_scale = dynamic_mxfp4_quant(A)
        A_scale = e8m0_shuffle(A_scale)
        A_q = A_q.view(dtypes.fp4x2)
        A_scale = A_scale.view(dtypes.fp8_e8m0)

    return aiter.gemm_a4w4(
        A_q, B_shuffle, A_scale, B_scale_sh,
        dtype=dtypes.bf16, bpreshuffle=True,
    )
scrolls · 188 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