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
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-k
split_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