Skip to content
KernelIndex
Search⌘K

submission 754473

wildman · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

amd-mxfp4-mm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-754473?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.3µs
#1043 of 1143
2026-04-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:70e5b0a36825806af4628bd78f30c324abd87aea3d91a3d35d1dbed783012a93
license declaredunknown
license concludedunknown
authorswildman
imported2026-08-26

Techniques

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

shared-memory__shared__ float mx[256];
split-ksplit_k = cfg.get("splitK", 0)

Kernel source

amd-mxfp4-mm.py360 lines
import functools
import hashlib
import os
import torch
from task import input_t, output_t

import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle

try:
    from aiter.ops.gemm_op_a4w4 import (
        get_GEMM_config as _get_GEMM_config,
        gemm_a4w4_asm as _gemm_a4w4_asm,
        gemm_a4w4_blockscale as _gemm_a4w4_blockscale,
    )
    _DIRECT = True
except Exception:
    _DIRECT = False

_dynq = dynamic_mxfp4_quant
_shuf = e8m0_shuffle
_wrap = aiter.gemm_a4w4
_fp4x2 = dtypes.fp4x2
_fp8e8m0 = dtypes.fp8_e8m0
_bf16 = dtypes.bf16

_MOD = None
_HWQ_OK = True
_BUF = {}

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


def _dev_id(device: torch.device) -> int:
    idx = device.index
    return torch.cuda.current_device() if idx is None else idx


def _shape_dims(m: int, k: int):
    gk = k >> 5
    sn = (gk + 7) & -8
    mpad = (m + 255) & -256
    qk = k >> 1
    return gk, sn, mpad, qk


def _bufs(device: torch.device, m: int, n: int, k: int):
    key = (_dev_id(device), m, n, k)
    buf = _BUF.get(key)
    if buf is not None:
        return buf

    _, sn, mpad, qk = _shape_dims(m, k)
    m32 = (m + 31) & -32

    q_u8 = torch.empty((m, qk), dtype=torch.uint8, device=device)
    s_u8 = torch.empty((mpad, sn), dtype=torch.uint8, device=device)
    out = torch.empty((m32, n), dtype=torch.bfloat16, device=device)

    buf = (
        q_u8,
        q_u8.view(_fp4x2),
        s_u8,
        s_u8.view(_fp8e8m0),
        out,
    )
    _BUF[key] = buf
    return buf


@functools.lru_cache(maxsize=256)
def _plan(m: int, n: int, k: int):
    if not _DIRECT:
        return 2, 0, ""

    cfg = _get_GEMM_config(m, n, k)
    if cfg is None:
        return 0, 0, ""

    split_k = cfg.get("splitK", 0)
    split_k = 0 if split_k is None else int(split_k)
    kernel_name = cfg["kernelName"]

    if "_ZN" not in kernel_name:
        return 1, split_k, ""
    return 0, split_k, kernel_name


def _prime_hot():
    if not _DIRECT:
        return
    for m, n, k in _HOT:
        _plan(m, n, k)


def _load_mod():
    global _MOD, _HWQ_OK
    if _MOD is not None or not _HWQ_OK:
        return _MOD
    if not (hasattr(torch.version, "hip") and torch.version.hip):
        _HWQ_OK = False
        return None

    os.environ.setdefault("CXX", "amdclang++")
    os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
    os.environ.setdefault("MAX_JOBS", "2")

    from torch.utils.cpp_extension import load_inline

    cpp = r"""
#include <torch/extension.h>
void quant_mxfp4_bf16_hw_sh_out(torch::Tensor x, torch::Tensor q, torch::Tensor s_sh);
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
  m.def("quant_mxfp4_bf16_hw_sh_out", &quant_mxfp4_bf16_hw_sh_out);
}
"""

    hip = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAGuard.h>
#include <hip/hip_runtime.h>
#include <hip/hip_bfloat16.h>
#include <stdint.h>

using torch::Tensor;

typedef __bf16 bf16_t;
typedef bf16_t v2bf16 __attribute__((ext_vector_type(2)));
typedef unsigned char u8;

__device__ __forceinline__ u8 f32_to_e8m0(float x) {
    union {
        float f;
        unsigned int u;
    } v;
    v.f = x;
    unsigned int exponent = (v.u >> 23) & 0xffu;
    if (exponent == 0xffu) return 0xffu;
    unsigned int round_case =
        ((v.u & 0x400000u) != 0u) &&
        (((v.u & 0x200000u) != 0u) || ((v.u & 0x1fffffu) != 0u) || (exponent > 0u));
    exponent += round_case;
    return (u8)exponent;
}

__device__ __forceinline__ float e8m0_to_f32(u8 x) {
    union {
        float f;
        unsigned int u;
    } v;
    v.u = (x == 0) ? 0x00400000u : (unsigned(x) << 23);
    return v.f;
}

__device__ __forceinline__ int sh_idx(int row, int col, int sn) {
    int d0 = sn >> 3;
    int a = row >> 5;
    int b = (row >> 4) & 1;
    int c = row & 15;
    int d = col >> 3;
    int e = (col >> 2) & 1;
    int f = col & 3;
    return ((((((a * d0) + d) * 4 + f) * 16 + c) * 2 + e) * 2 + b);
}

__global__ __launch_bounds__(256, 2) void quant_kernel16(
    const bf16_t* __restrict__ x,
    u8* __restrict__ q,
    u8* __restrict__ s_sh,
    int m,
    int k,
    int gk,
    int sn,
    int qkbytes
) {
    __shared__ float mx[256];
    __shared__ u8 sb[16];

    const int tid = int(threadIdx.x);
    const int grp = tid >> 4;
    const int lane = tid & 15;
    const int chunk = int(blockIdx.x) * 16 + grp;
    const int total = m * gk;
    const bool live = chunk < total;

    int row = 0;
    int gb = 0;
    v2bf16 src;
    src[0] = (__bf16)0;
    src[1] = (__bf16)0;
    float local_max = 0.0f;

    if (live) {
        row = chunk / gk;
        gb = chunk - row * gk;
        const v2bf16* p2 = reinterpret_cast<const v2bf16*>(x + row * k + gb * 32);
        src = p2[lane];
        float x0 = fabsf((float)src[0]);
        float x1 = fabsf((float)src[1]);
        local_max = x0 > x1 ? x0 : x1;
    }

    mx[tid] = local_max;
    __syncthreads();

    if (lane == 0 && live) {
        float amax = mx[grp * 16];
        #pragma unroll
        for (int i = 1; i < 16; ++i) {
            float v = mx[grp * 16 + i];
            amax = v > amax ? v : amax;
        }
        u8 s = f32_to_e8m0(amax * 0.16666667163372039794921875f);
        sb[grp] = s;
        s_sh[sh_idx(row, gb, sn)] = s;
    }

    __syncthreads();

    if (live) {
        float sf = e8m0_to_f32(sb[grp]);
        unsigned int pack = __builtin_amdgcn_cvt_scalef32_pk_fp4_bf16(0u, src, sf, 0);
        q[row * qkbytes + gb * 16 + lane] = (u8)pack;
    }
}

void quant_mxfp4_bf16_hw_sh_out(Tensor x, Tensor q, Tensor s_sh) {
    TORCH_CHECK(x.is_cuda(), "x must be cuda");
    TORCH_CHECK(q.is_cuda(), "q must be cuda");
    TORCH_CHECK(s_sh.is_cuda(), "s_sh must be cuda");
    TORCH_CHECK(x.scalar_type() == at::kBFloat16, "x must be bf16");
    TORCH_CHECK(q.scalar_type() == at::kByte, "q must be u8");
    TORCH_CHECK(s_sh.scalar_type() == at::kByte, "s_sh must be u8");
    TORCH_CHECK(x.dim() == 2, "x must be 2d");
    TORCH_CHECK(q.dim() == 2, "q must be 2d");
    TORCH_CHECK(s_sh.dim() == 2, "s_sh must be 2d");

    Tensor xc = x.contiguous();
    const int m = int(xc.size(0));
    const int k = int(xc.size(1));
    TORCH_CHECK((k & 63) == 0, "k must be divisible by 64");

    const int gk = k >> 5;
    const int sn = (gk + 7) & ~7;
    const int qk = k >> 1;

    TORCH_CHECK(int(q.size(0)) == m, "bad q rows");
    TORCH_CHECK(int(q.size(1)) == qk, "bad q cols");
    TORCH_CHECK(int(s_sh.size(0)) >= m, "bad s_sh rows");
    TORCH_CHECK(int(s_sh.size(1)) >= gk, "bad s_sh cols");
    TORCH_CHECK((int(s_sh.size(1)) & 7) == 0, "bad s_sh pad");

    at::cuda::CUDAGuard guard(xc.device());

    const int blocks = (m * gk + 15) >> 4;
    hipLaunchKernelGGL(
        quant_kernel16,
        dim3(blocks),
        dim3(256),
        0,
        0,
        reinterpret_cast<const bf16_t*>(xc.data_ptr<at::BFloat16>()),
        q.data_ptr<u8>(),
        s_sh.data_ptr<u8>(),
        m,
        k,
        gk,
        int(s_sh.size(1)),
        qk
    );

    hipError_t err = hipGetLastError();
    TORCH_CHECK(err == hipSuccess, hipGetErrorString(err));
}
"""

    key = hashlib.sha1((cpp + hip).encode()).hexdigest()[:10]

    try:
        _MOD = load_inline(
            name="mxfp4_hwq16_" + key,
            cpp_sources=[cpp],
            cuda_sources=[hip],
            functions=None,
            extra_cflags=["-O3", "-std=c++17"],
            extra_cuda_cflags=["-O3", "-std=c++17", "--offload-arch=gfx950"],
            with_cuda=True,
            verbose=False,
        )
        _prime_hot()
    except Exception:
        _HWQ_OK = False
        _MOD = None

    return _MOD


def _fallback_quant(a: torch.Tensor):
    q, s = _dynq(a.contiguous())
    return q.view(_fp4x2), _shuf(s).contiguous().view(_fp8e8m0)


def _run_gemm(aq, bw, asq, bs, m: int, n: int, k: int, device: torch.device):
    mode, split_k, kernel_name = _plan(m, n, k)
    if mode == 2:
        return _wrap(aq, bw, asq, bs, dtype=_bf16, bpreshuffle=True)

    out = _bufs(device, m, n, k)[4]

    if mode == 1:
        _gemm_a4w4_blockscale(aq, bw, asq, bs, out, splitK=split_k)
        return out[:m]

    _gemm_a4w4_asm(
        aq,
        bw,
        asq,
        bs,
        out,
        kernelName=kernel_name,
        bias=None,
        alpha=1.0,
        beta=0.0,
        bpreshuffle=True,
        log2_k_split=split_k,
    )
    return out[:m]


@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
    a = data[0]
    bw = data[3]
    bs = data[4]

    m, k = a.shape
    n = bw.shape[0]
    device = a.device

    mod = _load_mod()
    if mod is not None:
        q_u8, aq, s_u8, asq, _ = _bufs(device, m, n, k)
        mod.quant_mxfp4_bf16_hw_sh_out(a, q_u8, s_u8)
        return _run_gemm(aq, bw, asq, bs, m, n, k, device)

    aq, asq = _fallback_quant(a)
    return _run_gemm(aq, bw, asq, bs, m, n, k, device)
scrolls · 360 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