submission 620883
fchange · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 993 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-620883?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:20db7f3ed3970b66784e586f201c42b352b5e410cc2fda82f34bfbcce5fbe76b
license declaredunknown
license concludedunknown
authorsfchange
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MXFP4 per-1x32 quant on A, then A4W4 GEMM on MI355X.shared-memory
__shared__ float abs_vals[32];split-k
_A4W4_TUNED_HEADER = "cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio\n"stages = 1
NUM_STAGES=1,Kernel source
submission.py993 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
MXFP4 per-1x32 quant on A, then A4W4 GEMM on MI355X.
Formal submission path:
- exact-shape asm dispatch for the fixed benchmark shapes
- specialized quant+shuffle path for those same shapes
- unified aiter fallback for everything else
"""
import importlib.util
import os
from task import input_t, output_t
os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
os.environ.setdefault("CXX", "clang++")
_ASM_32X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_ASM_32X128_CO = "f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128.co"
_ASM_64X1024 = "_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_64x1024E"
_ASM_64X1024_CO = "f4gemm_bf16_per1x32Fp4_BpreShuffle_64x1024.co"
_ASM_64X512 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x512E"
_ASM_64X512_CO = "f4gemm_bf16_per1x32Fp4_BpreShuffle_64x512.co"
_ASM_128X512 = "_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_128x512E"
_ASM_128X512_CO = "f4gemm_bf16_per1x32Fp4_BpreShuffle_128x512.co"
_ASM_224X256 = "_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_224x256E"
_ASM_224X256_CO = "f4gemm_bf16_per1x32Fp4_BpreShuffle_224x256.co"
_ASM_256X256 = "_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_256x256E"
_ASM_256X256_CO = "f4gemm_bf16_per1x32Fp4_BpreShuffle_256x256.co"
_A4W4_TUNED_HEADER = "cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio\n"
_A4W4_TUNED_ROWS = [
(256, 4, 2880, 512, 21, 0, "0.1000", _ASM_32X128, "0.0", "0.0", "0.0"),
(256, 16, 2112, 7168, 21, 0, "0.1000", _ASM_32X128, "0.0", "0.0", "0.0"),
(256, 32, 4096, 512, 21, 0, "0.1000", _ASM_32X128, "0.0", "0.0", "0.0"),
(256, 32, 2880, 512, 21, 0, "0.1000", _ASM_32X128, "0.0", "0.0", "0.0"),
]
_SPECIALIZED_QUANT_ENV_VAR = "MXFP4_ENABLE_SPECIALIZED_QUANT"
_SPECIALIZED_QUANT_VARIANT_ENV_VAR = "MXFP4_SPECIALIZED_QUANT_VARIANT"
_LARGE_SHAPE_VARIANT_ENV_VAR = "MXFP4_LARGE_SHAPE_VARIANT"
_SPECIALIZED_QUANT_DEFAULT_VARIANT = "combo_m4m16m32cppqshuf"
_A4W4_TUNED_OVERRIDE = None
_UNIFIED_PLAN = {"kind": "unified"}
_SHAPE_PLANS = {
(4, 2880, 512): {
"kind": "asm",
"kernel_name": _ASM_32X128,
"co_name": _ASM_32X128_CO,
"log2_k_split": None,
"specialized_quant_candidate": True,
},
(16, 2112, 7168): {
"kind": "asm",
"kernel_name": _ASM_32X128,
"co_name": _ASM_32X128_CO,
"log2_k_split": None,
"specialized_quant_candidate": True,
},
(32, 4096, 512): {
"kind": "asm",
"kernel_name": _ASM_32X128,
"co_name": _ASM_32X128_CO,
"log2_k_split": None,
"specialized_quant_candidate": True,
},
(32, 2880, 512): {
"kind": "asm",
"kernel_name": _ASM_32X128,
"co_name": _ASM_32X128_CO,
"log2_k_split": None,
"specialized_quant_candidate": True,
},
(64, 7168, 2048): {
"kind": "unified",
"specialized_quant_candidate": True,
},
(256, 3072, 1536): {
"kind": "unified",
"specialized_quant_candidate": True,
},
}
_LARGE_SHAPE_VARIANTS = {
"baseline": {},
"64x1024_256x256s1": {
(64, 7168, 2048): {
"kind": "asm",
"kernel_name": _ASM_64X1024,
"co_name": _ASM_64X1024_CO,
"log2_k_split": None,
"specialized_quant_candidate": False,
},
(256, 3072, 1536): {
"kind": "asm",
"kernel_name": _ASM_256X256,
"co_name": _ASM_256X256_CO,
"log2_k_split": 1,
"specialized_quant_candidate": False,
},
},
"64x1024_256x256s2": {
(64, 7168, 2048): {
"kind": "asm",
"kernel_name": _ASM_64X1024,
"co_name": _ASM_64X1024_CO,
"log2_k_split": None,
"specialized_quant_candidate": False,
},
(256, 3072, 1536): {
"kind": "asm",
"kernel_name": _ASM_256X256,
"co_name": _ASM_256X256_CO,
"log2_k_split": 2,
"specialized_quant_candidate": False,
},
},
"128x512s1_both": {
(64, 7168, 2048): {
"kind": "asm",
"kernel_name": _ASM_128X512,
"co_name": _ASM_128X512_CO,
"log2_k_split": 1,
"specialized_quant_candidate": False,
},
(256, 3072, 1536): {
"kind": "asm",
"kernel_name": _ASM_128X512,
"co_name": _ASM_128X512_CO,
"log2_k_split": 1,
"specialized_quant_candidate": False,
},
},
"128x512s2_both": {
(64, 7168, 2048): {
"kind": "asm",
"kernel_name": _ASM_128X512,
"co_name": _ASM_128X512_CO,
"log2_k_split": 2,
"specialized_quant_candidate": False,
},
(256, 3072, 1536): {
"kind": "asm",
"kernel_name": _ASM_128X512,
"co_name": _ASM_128X512_CO,
"log2_k_split": 2,
"specialized_quant_candidate": False,
},
},
}
_RUNTIME = None
_SPECIALIZED_QUANT_RUNTIME = None
_NATIVE_SHUFFLE_RUNTIME = None
_SPECIALIZED_QUANT_WORKSPACES = {}
_SPECIALIZED_QUANT_ERROR = None
_SPECIALIZED_QUANT_INFO_PRINTED = False
_NATIVE_QUANT_ERROR = None
_NATIVE_QUANT_INFO_PRINTED = False
_NATIVE_SHUFFLE_ERROR = None
_NATIVE_SHUFFLE_INFO_PRINTED = False
def _select_gemm_plan(m: int, n: int, k: int):
shape = (m, n, k)
variant = os.environ.get(_LARGE_SHAPE_VARIANT_ENV_VAR, "baseline")
variant_plans = _LARGE_SHAPE_VARIANTS.get(variant, _LARGE_SHAPE_VARIANTS["baseline"])
if shape in variant_plans:
return variant_plans[shape]
return _SHAPE_PLANS.get(shape, _UNIFIED_PLAN)
def _find_aiter_config_path():
try:
spec = importlib.util.find_spec("aiter")
except (ImportError, ValueError):
spec = None
locations = getattr(spec, "submodule_search_locations", None) if spec is not None else None
if locations:
return os.path.join(locations[0], "configs", "a4w4_blockscale_tuned_gemm.csv")
runtime = globals().get("_RUNTIME")
if runtime is None:
return None
_, aiter, _, _, _ = runtime
package_file = getattr(aiter, "__file__", None)
if package_file:
return os.path.join(
os.path.dirname(os.path.abspath(package_file)),
"configs",
"a4w4_blockscale_tuned_gemm.csv",
)
return None
def _render_a4w4_tuned_override():
rows = ["{},{},{},{},{},{},{},{},{},{},{}".format(*row) for row in _A4W4_TUNED_ROWS]
return _A4W4_TUNED_HEADER + "\n".join(rows) + "\n"
def _ensure_a4w4_tuned_override():
global _A4W4_TUNED_OVERRIDE
override_path = _A4W4_TUNED_OVERRIDE or "/tmp/mxfp4_a4w4_tuned_override.csv"
content = _render_a4w4_tuned_override()
try:
existing = None
if os.path.exists(override_path):
with open(override_path, "r", encoding="utf-8") as handle:
existing = handle.read()
if existing != content:
with open(override_path, "w", encoding="utf-8") as handle:
handle.write(content)
except Exception:
return None
_A4W4_TUNED_OVERRIDE = override_path
default_path = _find_aiter_config_path()
if not default_path:
return override_path
os.environ["AITER_CONFIG_GEMM_A4W4"] = os.pathsep.join([default_path, override_path])
return override_path
def _get_runtime():
global _RUNTIME
if _RUNTIME is None:
_ensure_a4w4_tuned_override()
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
_RUNTIME = (torch, aiter, dtypes, dynamic_mxfp4_quant, e8m0_shuffle)
return _RUNTIME
def _alloc_gemm_output(a_q, dtypes, m: int, n: int, zero_init: bool = False):
out_rows = ((m + 31) // 32) * 32
if zero_init:
return a_q.new_zeros((out_rows, n), dtype=dtypes.bf16)
return a_q.new_empty((out_rows, n), dtype=dtypes.bf16)
def _quant_mxfp4(x, dtypes, dynamic_mxfp4_quant, e8m0_shuffle):
x_fp4, scale = dynamic_mxfp4_quant(x)
scale = e8m0_shuffle(scale)
return x_fp4.view(dtypes.fp4x2), scale.view(dtypes.fp8_e8m0)
def _specialized_quant_runtime_enabled(torch, plan):
if not plan.get("specialized_quant_candidate"):
return False
if os.environ.get(_SPECIALIZED_QUANT_ENV_VAR, "1") == "0":
return False
return getattr(getattr(torch, "version", None), "hip", None) is not None
def _get_specialized_quant_runtime():
global _SPECIALIZED_QUANT_RUNTIME
if _SPECIALIZED_QUANT_RUNTIME is not None:
return _SPECIALIZED_QUANT_RUNTIME
import triton
import triton.language as tl
from aiter.ops.triton._triton_kernels.quant.quant import _dynamic_mxfp4_quant_kernel
@triton.jit
def _shuffle_e8m0_scale_kernel(
src_ptr,
dst_ptr,
stride_src_m,
stride_src_n,
M,
N_VALID,
N_PAD,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
mask = (offs_m[:, None] < M) & (offs_n[None, :] < N_VALID)
src_offs = offs_m[:, None] * stride_src_m + offs_n[None, :] * stride_src_n
vals = tl.load(src_ptr + src_offs, mask=mask, other=127)
g0 = offs_m[:, None] // 32
rem_m = offs_m[:, None] % 32
g1 = rem_m // 16
g2 = rem_m % 16
g3 = offs_n[None, :] // 8
rem_n = offs_n[None, :] % 8
g4 = rem_n // 4
g5 = rem_n % 4
dst_offs = (
g1
+ g4 * 2
+ g2 * 4
+ g5 * 64
+ g3 * 256
+ g0 * (32 * N_PAD)
)
tl.store(dst_ptr + dst_offs, vals, mask=mask)
_SPECIALIZED_QUANT_RUNTIME = (triton, _dynamic_mxfp4_quant_kernel, _shuffle_e8m0_scale_kernel)
return _SPECIALIZED_QUANT_RUNTIME
def _get_specialized_quant_workspace(torch, x):
m, n = x.shape
scale_n_valid = (n + 31) // 32
scale_n_pad = ((scale_n_valid + 7) // 8) * 8
scale_m_pad = ((m + 255) // 256) * 256
cache_key = (tuple(x.shape), str(getattr(x, "device", "")), str(getattr(x, "dtype", "")))
workspace = _SPECIALIZED_QUANT_WORKSPACES.get(cache_key)
if workspace is None:
workspace = {
"a_q_raw": torch.empty((m, n // 2), dtype=torch.uint8, device=x.device),
"a_scale_raw": torch.empty((m, scale_n_valid), dtype=torch.uint8, device=x.device),
"a_scale_shuffled_raw": torch.full(
(scale_m_pad, scale_n_pad),
127,
dtype=torch.uint8,
device=x.device,
),
"scale_n_valid": scale_n_valid,
"scale_n_pad": scale_n_pad,
}
_SPECIALIZED_QUANT_WORKSPACES[cache_key] = workspace
return workspace
def _get_specialized_quant_launch_config(triton, m: int, n: int):
variant = os.environ.get(_SPECIALIZED_QUANT_VARIANT_ENV_VAR, _SPECIALIZED_QUANT_DEFAULT_VARIANT)
if variant == "combo_best" and (m, n) == (16, 7168):
return {
"num_iter": 1,
"block_size_m": 16,
"block_size_n": 256,
"num_warps": 4,
}
if variant == "combo_best" and (m, n) in ((64, 2048), (256, 1536)):
return {
"num_iter": 2,
"block_size_m": 64,
"block_size_n": 128,
"num_warps": 4,
}
if variant == "combo_m16sh16" and (m, n) == (16, 7168):
return {
"num_iter": 1,
"block_size_m": 16,
"block_size_n": 256,
"num_warps": 4,
}
if variant == "combo_m16sh16" and (m, n) in ((64, 2048), (256, 1536)):
return {
"num_iter": 2,
"block_size_m": 64,
"block_size_n": 128,
"num_warps": 4,
}
if variant == "combo_m16n128sh16" and (m, n) == (16, 7168):
return {
"num_iter": 1,
"block_size_m": 16,
"block_size_n": 128,
"num_warps": 4,
}
if variant == "combo_m16n128sh16" and (m, n) in ((64, 2048), (256, 1536)):
return {
"num_iter": 2,
"block_size_m": 64,
"block_size_n": 128,
"num_warps": 4,
}
if variant == "combo_m16n128cppshuf" and (m, n) == (16, 7168):
return {
"num_iter": 1,
"block_size_m": 16,
"block_size_n": 128,
"num_warps": 4,
}
if variant == "combo_m16n128cppshuf" and (m, n) in ((64, 2048), (256, 1536)):
return {
"num_iter": 2,
"block_size_m": 64,
"block_size_n": 128,
"num_warps": 4,
}
if variant == "combo_m16n128cppqshuf" and (m, n) == (16, 7168):
return {
"num_iter": 1,
"block_size_m": 16,
"block_size_n": 128,
"num_warps": 4,
}
if variant == "combo_m16n128cppqshuf" and (m, n) in ((64, 2048), (256, 1536)):
return {
"num_iter": 2,
"block_size_m": 64,
"block_size_n": 128,
"num_warps": 4,
}
if variant == "combo_m16m32cppqshuf" and (m, n) == (16, 7168):
return {
"num_iter": 1,
"block_size_m": 16,
"block_size_n": 128,
"num_warps": 4,
}
if variant == "combo_m16m32cppqshuf" and (m, n) in ((64, 2048), (256, 1536)):
return {
"num_iter": 2,
"block_size_m": 64,
"block_size_n": 128,
"num_warps": 4,
}
if variant == "combo_m4m16m32cppqshuf" and (m, n) == (16, 7168):
return {
"num_iter": 1,
"block_size_m": 16,
"block_size_n": 128,
"num_warps": 4,
}
if variant == "combo_m4m16m32cppqshuf" and (m, n) in ((64, 2048), (256, 1536)):
return {
"num_iter": 2,
"block_size_m": 64,
"block_size_n": 128,
"num_warps": 4,
}
if variant == "lg64x64" and (m, n) in ((64, 2048), (256, 1536)):
return {
"num_iter": 4,
"block_size_m": 64,
"block_size_n": 64,
"num_warps": 4,
}
if variant == "lg64x128i2" and (m, n) in ((64, 2048), (256, 1536)):
return {
"num_iter": 2,
"block_size_m": 64,
"block_size_n": 128,
"num_warps": 4,
}
if variant == "lg32x128w8" and (m, n) in ((64, 2048), (256, 1536)):
return {
"num_iter": 4,
"block_size_m": 32,
"block_size_n": 128,
"num_warps": 8,
}
if variant == "m16n128w4" and (m, n) == (16, 7168):
return {
"num_iter": 1,
"block_size_m": 16,
"block_size_n": 128,
"num_warps": 4,
}
if variant == "m16n256w4" and (m, n) == (16, 7168):
return {
"num_iter": 1,
"block_size_m": 16,
"block_size_n": 256,
"num_warps": 4,
}
if m <= 32:
return {
"num_iter": 1,
"block_size_m": triton.next_power_of_2(m),
"block_size_n": 32,
"num_warps": 1,
}
config = {
"num_iter": 4,
"block_size_m": 64,
"block_size_n": 64,
"num_warps": 4,
}
if n <= 16384:
config["block_size_m"] = 32
config["block_size_n"] = 128
if n <= 1024:
config["num_iter"] = 1
config["block_size_n"] = min(256, triton.next_power_of_2(n))
config["block_size_n"] = max(32, config["block_size_n"])
config["block_size_m"] = min(8, triton.next_power_of_2(m))
config["num_warps"] = 4
return config
def _get_specialized_shuffle_launch_config(m: int, scale_n_valid: int):
variant = os.environ.get(_SPECIALIZED_QUANT_VARIANT_ENV_VAR, _SPECIALIZED_QUANT_DEFAULT_VARIANT)
if variant in ("m16sh16", "combo_m16sh16", "combo_m16n128sh16") and (m, scale_n_valid) == (16, 224):
return {
"block_m": 16,
"block_n": 16,
"num_warps": 1,
}
return {
"block_m": 32,
"block_n": 8,
"num_warps": 1,
}
def _native_shuffle_enabled(m: int, scale_n_valid: int):
variant = os.environ.get(_SPECIALIZED_QUANT_VARIANT_ENV_VAR, _SPECIALIZED_QUANT_DEFAULT_VARIANT)
return variant == "combo_m16n128cppshuf" and (m, scale_n_valid) == (16, 224)
def _native_quant_enabled(m: int, n: int):
variant = os.environ.get(_SPECIALIZED_QUANT_VARIANT_ENV_VAR, _SPECIALIZED_QUANT_DEFAULT_VARIANT)
if variant == "combo_m16n128cppqshuf":
return (m, n) == (16, 7168)
if variant == "combo_m16m32cppqshuf":
return (m, n) in ((16, 7168), (32, 512))
if variant == "combo_m4m16m32cppqshuf":
return (m, n) in ((4, 512), (16, 7168), (32, 512))
return False
def _get_native_shuffle_module(torch):
global _NATIVE_SHUFFLE_RUNTIME
if _NATIVE_SHUFFLE_RUNTIME is not None:
return _NATIVE_SHUFFLE_RUNTIME
from torch.utils.cpp_extension import load_inline
cpp_src = """
void mxfp4_quant_shuffle_small(
torch::Tensor x,
torch::Tensor x_fp4,
torch::Tensor bs_shuffled,
int64_t m,
int64_t n,
int64_t n_pad);
void mxfp4_shuffle_e8m0(
torch::Tensor src,
torch::Tensor dst,
int64_t m,
int64_t n_valid,
int64_t n_pad);
"""
hip_src = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>
#include <cstdint>
#include <cmath>
__device__ __forceinline__ uint8_t float_to_e8m0(float x) {
if (x <= 0.0f) {
return 0;
}
uint32_t u = __float_as_uint(x);
uint32_t exponent = (u >> 23) & 0xFF;
bool round_case = ((u & 0x400000) > 0) &&
(((u & 0x200000) > 0) || ((u & 0x1FFFFF) > 0) || (exponent > 0));
if (round_case && exponent < 0xFF) {
exponent += 1;
}
return static_cast<uint8_t>(exponent);
}
__device__ __forceinline__ float e8m0_to_float(uint8_t x) {
if (x == 0) {
return __uint_as_float(0x00400000u);
}
if (x == 0xFF) {
return __uint_as_float(0x7F800001u);
}
return __uint_as_float(static_cast<uint32_t>(x) << 23);
}
__device__ __forceinline__ uint8_t float_to_mxfp4(float x) {
constexpr int EXP_BIAS_FP32 = 127;
constexpr int EXP_BIAS_FP4 = 1;
constexpr int EBITS_F32 = 8;
constexpr int EBITS_FP4 = 2;
constexpr int MBITS_F32 = 23;
constexpr int MBITS_FP4 = 1;
constexpr float MAX_NORMAL = 6.0f;
constexpr float MIN_NORMAL = 1.0f;
constexpr uint8_t MAX_INT = 0x7;
uint32_t qx = __float_as_uint(x);
uint32_t sign = qx & 0x80000000u;
qx ^= sign;
float qx_fp32 = __uint_as_float(qx);
bool saturate_mask = qx_fp32 >= MAX_NORMAL;
bool denormal_mask = (!saturate_mask) && (qx_fp32 < MIN_NORMAL);
bool normal_mask = !(saturate_mask || denormal_mask);
constexpr int denorm_exp = (EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1;
constexpr uint32_t denorm_mask_int = static_cast<uint32_t>(denorm_exp) << MBITS_F32;
float denorm_mask_float = __uint_as_float(denorm_mask_int);
uint8_t denormal_x = static_cast<uint8_t>(__float_as_uint(qx_fp32 + denorm_mask_float) - denorm_mask_int);
uint32_t normal_x = qx;
uint32_t mant_odd = (normal_x >> (MBITS_F32 - MBITS_FP4)) & 1u;
constexpr int32_t val_to_add = ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1 << 21) - 1;
normal_x = static_cast<uint32_t>(static_cast<int32_t>(normal_x) + val_to_add);
normal_x += mant_odd;
normal_x = normal_x >> (MBITS_F32 - MBITS_FP4);
uint8_t out = MAX_INT;
if (normal_mask) {
out = static_cast<uint8_t>(normal_x);
}
if (denormal_mask) {
out = denormal_x;
}
uint8_t sign_lp = static_cast<uint8_t>(sign >> (MBITS_F32 + EBITS_F32 - MBITS_FP4 - EBITS_FP4));
return out | sign_lp;
}
__global__ void mxfp4_quant_shuffle_small_kernel(
const __hip_bfloat16* x,
uint8_t* x_fp4,
uint8_t* bs_shuffled,
int64_t stride_x_m,
int64_t stride_x_n,
int64_t stride_x_fp4_m,
int64_t stride_x_fp4_n,
int64_t m,
int64_t n,
int64_t n_pad) {
__shared__ float abs_vals[32];
__shared__ uint8_t fp4_codes[32];
__shared__ uint8_t bs_e8m0;
int lane = threadIdx.x;
int row = blockIdx.y;
int block_n = blockIdx.x;
int col = block_n * 32 + lane;
if (row >= m || col >= n) {
return;
}
float x_val = __bfloat162float(x[row * stride_x_m + col * stride_x_n]);
abs_vals[lane] = fabsf(x_val);
__syncthreads();
for (int offset = 16; offset > 0; offset >>= 1) {
if (lane < offset) {
abs_vals[lane] = fmaxf(abs_vals[lane], abs_vals[lane + offset]);
}
__syncthreads();
}
if (lane == 0) {
float amax = abs_vals[0];
if (amax == 0.0f) {
bs_e8m0 = 0;
} else {
uint32_t amax_bits = __float_as_uint(amax);
amax_bits = (amax_bits + 0x200000u) & 0xFF800000u;
float rounded_amax = __uint_as_float(amax_bits);
bs_e8m0 = float_to_e8m0(rounded_amax / 4.0f);
}
}
__syncthreads();
float scale = e8m0_to_float(bs_e8m0);
float qx = x_val / scale;
fp4_codes[lane] = float_to_mxfp4(qx);
__syncthreads();
if (lane < 16) {
uint8_t even = fp4_codes[lane * 2];
uint8_t odd = fp4_codes[lane * 2 + 1];
x_fp4[row * stride_x_fp4_m + (block_n * 16 + lane) * stride_x_fp4_n] = static_cast<uint8_t>(even | (odd << 4));
}
if (lane == 0) {
int64_t g1 = row / 16;
int64_t g2 = row % 16;
int64_t g3 = block_n / 8;
int64_t rem_n = block_n % 8;
int64_t g4 = rem_n / 4;
int64_t g5 = rem_n % 4;
int64_t dst_off = g1 + g4 * 2 + g2 * 4 + g5 * 64 + g3 * 256;
bs_shuffled[dst_off] = bs_e8m0;
}
}
__global__ void mxfp4_shuffle_e8m0_kernel(
const uint8_t* src,
uint8_t* dst,
int64_t stride_src_m,
int64_t stride_src_n,
int64_t m,
int64_t n_valid,
int64_t n_pad) {
int row = blockIdx.y * blockDim.y + threadIdx.y;
int col = blockIdx.x * blockDim.x + threadIdx.x;
if (row >= m || col >= n_valid) {
return;
}
uint8_t val = src[row * stride_src_m + col * stride_src_n];
int64_t g0 = row / 32;
int64_t rem_m = row % 32;
int64_t g1 = rem_m / 16;
int64_t g2 = rem_m % 16;
int64_t g3 = col / 8;
int64_t rem_n = col % 8;
int64_t g4 = rem_n / 4;
int64_t g5 = rem_n % 4;
int64_t dst_off =
g1 + g4 * 2 + g2 * 4 + g5 * 64 + g3 * 256 + g0 * (32 * n_pad);
dst[dst_off] = val;
}
void mxfp4_shuffle_e8m0(
torch::Tensor src,
torch::Tensor dst,
int64_t m,
int64_t n_valid,
int64_t n_pad) {
TORCH_CHECK(src.is_cuda(), "src must be CUDA/HIP");
TORCH_CHECK(dst.is_cuda(), "dst must be CUDA/HIP");
TORCH_CHECK(src.scalar_type() == torch::kUInt8, "src must be uint8");
TORCH_CHECK(dst.scalar_type() == torch::kUInt8, "dst must be uint8");
const dim3 threads(16, 16);
const dim3 blocks(
static_cast<unsigned int>((n_valid + threads.x - 1) / threads.x),
static_cast<unsigned int>((m + threads.y - 1) / threads.y));
hipLaunchKernelGGL(
mxfp4_shuffle_e8m0_kernel,
blocks,
threads,
0,
0,
src.data_ptr<uint8_t>(),
dst.data_ptr<uint8_t>(),
src.stride(0),
src.stride(1),
m,
n_valid,
n_pad);
auto err = hipGetLastError();
TORCH_CHECK(err == hipSuccess, hipGetErrorString(err));
}
void mxfp4_quant_shuffle_small(
torch::Tensor x,
torch::Tensor x_fp4,
torch::Tensor bs_shuffled,
int64_t m,
int64_t n,
int64_t n_pad) {
TORCH_CHECK(x.is_cuda(), "x must be CUDA/HIP");
TORCH_CHECK(x.scalar_type() == torch::kBFloat16, "x must be bf16");
TORCH_CHECK(x_fp4.scalar_type() == torch::kUInt8, "x_fp4 must be uint8");
TORCH_CHECK(bs_shuffled.scalar_type() == torch::kUInt8, "bs_shuffled must be uint8");
TORCH_CHECK(m == 4 || m == 16 || m == 32, "native quant probe only supports m=4, m=16, or m=32");
TORCH_CHECK((n % 32) == 0, "n must be divisible by 32");
const dim3 threads(32);
const dim3 blocks(
static_cast<unsigned int>(n / 32),
static_cast<unsigned int>(m));
hipLaunchKernelGGL(
mxfp4_quant_shuffle_small_kernel,
blocks,
threads,
0,
0,
reinterpret_cast<const __hip_bfloat16*>(x.data_ptr()),
x_fp4.data_ptr<uint8_t>(),
bs_shuffled.data_ptr<uint8_t>(),
x.stride(0),
x.stride(1),
x_fp4.stride(0),
x_fp4.stride(1),
m,
n,
n_pad);
auto err = hipGetLastError();
TORCH_CHECK(err == hipSuccess, hipGetErrorString(err));
}
"""
_NATIVE_SHUFFLE_RUNTIME = load_inline(
name="mxfp4_native_shuffle_v1",
cpp_sources=[cpp_src],
cuda_sources=[hip_src],
functions=["mxfp4_quant_shuffle_small", "mxfp4_shuffle_e8m0"],
extra_cflags=["-std=c++20"],
extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20"],
verbose=False,
)
return _NATIVE_SHUFFLE_RUNTIME
def _maybe_native_shuffle(torch, src, dst, m: int, n_valid: int, n_pad: int):
global _NATIVE_SHUFFLE_ERROR, _NATIVE_SHUFFLE_INFO_PRINTED
if not _native_shuffle_enabled(m, n_valid):
return False
try:
module = _get_native_shuffle_module(torch)
module.mxfp4_shuffle_e8m0(src, dst, m, n_valid, n_pad)
except Exception:
if _NATIVE_SHUFFLE_ERROR is None:
_NATIVE_SHUFFLE_ERROR = True
try:
import traceback
print("[mxfp4 native shuffle] falling back after error:")
traceback.print_exc()
except Exception:
pass
return False
if not _NATIVE_SHUFFLE_INFO_PRINTED:
_NATIVE_SHUFFLE_INFO_PRINTED = True
try:
print("[mxfp4 native shuffle] using load_inline HIP path")
except Exception:
pass
return True
def _maybe_native_quant_and_shuffle(torch, x, a_q_raw, a_scale_shuffled_raw, m: int, n: int, n_pad: int):
global _NATIVE_QUANT_ERROR, _NATIVE_QUANT_INFO_PRINTED
if not _native_quant_enabled(m, n):
return False
try:
module = _get_native_shuffle_module(torch)
module.mxfp4_quant_shuffle_small(x, a_q_raw, a_scale_shuffled_raw, m, n, n_pad)
except Exception:
if _NATIVE_QUANT_ERROR is None:
_NATIVE_QUANT_ERROR = True
try:
import traceback
print("[mxfp4 native quant] falling back after error:")
traceback.print_exc()
except Exception:
pass
return False
if not _NATIVE_QUANT_INFO_PRINTED:
_NATIVE_QUANT_INFO_PRINTED = True
try:
print("[mxfp4 native quant] using load_inline HIP path")
except Exception:
pass
return True
def _quant_mxfp4_specialized(torch, dtypes, x):
triton, quant_kernel, shuffle_kernel = _get_specialized_quant_runtime()
workspace = _get_specialized_quant_workspace(torch, x)
m, n = x.shape
a_q_raw = workspace["a_q_raw"]
a_scale_raw = workspace["a_scale_raw"]
a_scale_shuffled_raw = workspace["a_scale_shuffled_raw"]
scale_n_valid = workspace["scale_n_valid"]
scale_n_pad = workspace["scale_n_pad"]
native_quant_used = _maybe_native_quant_and_shuffle(torch, x, a_q_raw, a_scale_shuffled_raw, m, n, scale_n_pad)
if not native_quant_used:
launch = _get_specialized_quant_launch_config(triton, m, n)
num_iter = launch["num_iter"]
block_size_m = launch["block_size_m"]
block_size_n = launch["block_size_n"]
num_warps = launch["num_warps"]
grid = (triton.cdiv(m, block_size_m), triton.cdiv(n, block_size_n * num_iter))
quant_kernel[grid](
x,
a_q_raw,
a_scale_raw,
*x.stride(),
*a_q_raw.stride(),
*a_scale_raw.stride(),
M=m,
N=n,
MXFP4_QUANT_BLOCK_SIZE=32,
SCALING_MODE=0,
NUM_ITER=num_iter,
BLOCK_SIZE_M=block_size_m,
BLOCK_SIZE_N=block_size_n,
NUM_STAGES=1,
num_warps=num_warps,
waves_per_eu=0,
num_stages=1,
)
if not native_quant_used and not _maybe_native_shuffle(torch, a_scale_raw, a_scale_shuffled_raw, m, scale_n_valid, scale_n_pad):
shuffle_launch = _get_specialized_shuffle_launch_config(m, scale_n_valid)
shuffle_block_m = shuffle_launch["block_m"]
shuffle_block_n = shuffle_launch["block_n"]
shuffle_grid = (triton.cdiv(m, shuffle_block_m), triton.cdiv(scale_n_valid, shuffle_block_n))
shuffle_kernel[shuffle_grid](
a_scale_raw,
a_scale_shuffled_raw,
*a_scale_raw.stride(),
M=m,
N_VALID=scale_n_valid,
N_PAD=scale_n_pad,
BLOCK_M=shuffle_block_m,
BLOCK_N=shuffle_block_n,
num_warps=shuffle_launch["num_warps"],
num_stages=1,
)
return a_q_raw.view(dtypes.fp4x2), a_scale_shuffled_raw.view(dtypes.fp8_e8m0)
def _maybe_quant_mxfp4_specialized(torch, dtypes, x, plan):
global _SPECIALIZED_QUANT_ERROR, _SPECIALIZED_QUANT_INFO_PRINTED
if not _specialized_quant_runtime_enabled(torch, plan):
return None
try:
result = _quant_mxfp4_specialized(torch, dtypes, x)
except Exception:
if _SPECIALIZED_QUANT_ERROR is None:
_SPECIALIZED_QUANT_ERROR = True
try:
import traceback
print("[mxfp4 quant] falling back after error:")
traceback.print_exc()
except Exception:
pass
return None
if not _SPECIALIZED_QUANT_INFO_PRINTED:
_SPECIALIZED_QUANT_INFO_PRINTED = True
try:
print("[mxfp4 quant] using specialized quant+shuffle path")
except Exception:
pass
return result
def _run_gemm_asm(aiter, dtypes, a_q, b_shuffle, a_scale_sh, b_scale_sh, m: int, n: int, plan):
out = _alloc_gemm_output(a_q, dtypes, m, n)
aiter.gemm_a4w4_asm(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
out,
plan["kernel_name"],
bpreshuffle=True,
log2_k_split=plan["log2_k_split"],
)
return out[:m]
def custom_kernel(data: input_t) -> output_t:
torch, aiter, dtypes, dynamic_mxfp4_quant, e8m0_shuffle = _get_runtime()
a, b, _b_q, b_shuffle, b_scale_sh = data
a = a.contiguous()
b = b.contiguous()
m, k = a.shape
n, _ = b.shape
plan = _select_gemm_plan(m, n, k)
specialized = _maybe_quant_mxfp4_specialized(torch, dtypes, a, plan)
if specialized is None:
a_q, a_scale_sh = _quant_mxfp4(a, dtypes, dynamic_mxfp4_quant, e8m0_shuffle)
else:
a_q, a_scale_sh = specialized
if plan.get("kind") == "asm":
return _run_gemm_asm(
aiter,
dtypes,
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
m,
n,
plan,
)
return aiter.gemm_a4w4(
a_q,
b_shuffle,
a_scale_sh,
b_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
scrolls · 993 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