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
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.
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