submission 629301
francochengcc_30299 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1097 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-629301?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:a67d31c6dd6d7bed9f0aa94a5605b66c0a7aac9cee87f386d7fcd788dfba64fa
license declaredunknown
license concludedunknown
authorsfrancochengcc_30299
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[64];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.py1097 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"
_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"
_SMALLM_TRITON_ENV_VAR = "MXFP4_ENABLE_SMALLM_TRITON"
_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,
},
}
_NATIVE_QUANT_SHAPES = {
(4, 512),
(16, 7168),
(32, 512),
(64, 2048),
(256, 1536),
}
_DIRECT_QUANT_CONFIGS = {
(16, 7168): {
"num_iter": 1,
"block_size_m": 16,
"block_size_n": 128,
"num_warps": 4,
},
(64, 2048): {
"num_iter": 2,
"block_size_m": 64,
"block_size_n": 128,
"num_warps": 4,
},
(256, 1536): {
"num_iter": 2,
"block_size_m": 64,
"block_size_n": 128,
"num_warps": 4,
},
}
_RUNTIME = None
_SPECIALIZED_QUANT_RUNTIME = None
_NATIVE_SHUFFLE_RUNTIME = None
_SMALLM_TRITON_RUNTIME = None
_SPECIALIZED_QUANT_WORKSPACES = {}
_GEMM_OUTPUT_WORKSPACES = {}
_SMALLM_TRITON_WEIGHT_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
_SMALLM_TRITON_ERROR = None
_SMALLM_TRITON_INFO_PRINTED = False
def _select_gemm_plan(m: int, n: int, k: int):
return _SHAPE_PLANS.get((m, n, k), _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)
cache_key = ((out_rows, n), str(getattr(a_q, "device", "")), str(getattr(dtypes, "bf16", "bf16")))
out = _GEMM_OUTPUT_WORKSPACES.get(cache_key)
if out is None:
out = a_q.new_empty((out_rows, n), dtype=dtypes.bf16)
_GEMM_OUTPUT_WORKSPACES[cache_key] = out
return out
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 + 31) // 32) * 32
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):
direct = _DIRECT_QUANT_CONFIGS.get((m, n))
if direct is not None:
return direct
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):
return {
"block_m": 32,
"block_n": 8,
"num_warps": 1,
}
def _native_shuffle_enabled(m: int, scale_n_valid: int):
del m, scale_n_valid
return False
def _native_quant_enabled(m: int, n: int):
return (m, n) in _NATIVE_QUANT_SHAPES
def _smallm_triton_enabled(torch, m: int, n: int, k: int, plan):
if os.environ.get(_SMALLM_TRITON_ENV_VAR, "0") == "0":
return False
if plan.get("kind") != "asm":
return False
if (m, n, k) != (16, 2112, 7168):
return False
if getattr(getattr(torch, "version", None), "hip", None) is None:
return False
return hasattr(torch, "Tensor")
def _get_smallm_triton_runtime():
global _SMALLM_TRITON_RUNTIME
if _SMALLM_TRITON_RUNTIME is not None:
return _SMALLM_TRITON_RUNTIME
import triton
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (
_gemm_a16wfp4_preshuffle_kernel,
_get_config,
)
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
_gemm_afp4wfp4_reduce_kernel,
)
_SMALLM_TRITON_RUNTIME = (
triton,
_gemm_a16wfp4_preshuffle_kernel,
_gemm_afp4wfp4_reduce_kernel,
_get_config,
)
return _SMALLM_TRITON_RUNTIME
def _run_smallm_triton_preshuffle(torch, dtypes, x, b_shuffle, b_scale_sh, m: int, n: int):
triton, preshuffle_kernel, reduce_kernel, get_config = _get_smallm_triton_runtime()
b_shuffle_u8, b_scale_sh_u8 = _get_smallm_triton_weight_views(torch, b_shuffle, b_scale_sh, n)
packed_k = int(b_shuffle_u8.shape[1] // 16)
config, _ = get_config(m, n, packed_k, True)
config = dict(config)
num_ksplit = int(config.get("NUM_KSPLIT", 1))
block_size_k = int(config["BLOCK_SIZE_K"])
config["BLOCK_SIZE_N"] = max(32, int(config["BLOCK_SIZE_N"]))
if block_size_k >= 2 * packed_k:
block_size_k = triton.next_power_of_2(2 * packed_k)
config["BLOCK_SIZE_K"] = block_size_k
config["SPLITK_BLOCK_SIZE"] = 2 * packed_k
config["NUM_KSPLIT"] = 1
num_ksplit = 1
else:
config["SPLITK_BLOCK_SIZE"] = 2 * packed_k
y = torch.empty((m, n), dtype=dtypes.bf16, device=x.device)
y_pp = None
if num_ksplit > 1:
y_pp = torch.empty((num_ksplit, m, n), dtype=torch.float32, device=x.device)
grid = (
config["NUM_KSPLIT"]
* triton.cdiv(m, config["BLOCK_SIZE_M"])
* triton.cdiv(n, config["BLOCK_SIZE_N"]),
)
preshuffle_kernel[grid](
x,
b_shuffle_u8,
y if config["NUM_KSPLIT"] == 1 else y_pp,
b_scale_sh_u8,
m,
n,
packed_k,
x.stride(0),
x.stride(1),
b_shuffle_u8.stride(0),
b_shuffle_u8.stride(1),
0 if config["NUM_KSPLIT"] == 1 else y_pp.stride(0),
y.stride(0) if config["NUM_KSPLIT"] == 1 else y_pp.stride(1),
y.stride(1) if config["NUM_KSPLIT"] == 1 else y_pp.stride(2),
b_scale_sh_u8.stride(0),
b_scale_sh_u8.stride(1),
PREQUANT=True,
**config,
)
if config["NUM_KSPLIT"] > 1:
reduce_block_m = 16
reduce_block_n = 64
actual_ksplit = triton.cdiv(packed_k, (config["SPLITK_BLOCK_SIZE"] // 2))
grid_reduce = (
triton.cdiv(m, reduce_block_m),
triton.cdiv(n, reduce_block_n),
)
reduce_kernel[grid_reduce](
y_pp,
y,
m,
n,
y_pp.stride(0),
y_pp.stride(1),
y_pp.stride(2),
y.stride(0),
y.stride(1),
reduce_block_m,
reduce_block_n,
actual_ksplit,
triton.next_power_of_2(config["NUM_KSPLIT"]),
)
return y
def _get_smallm_triton_weight_views(torch, b_shuffle, b_scale_sh, n: int):
key = (
int(getattr(b_shuffle, "data_ptr", lambda: 0)()),
int(getattr(b_scale_sh, "data_ptr", lambda: 0)()),
tuple(b_shuffle.shape),
tuple(b_scale_sh.shape),
str(getattr(b_shuffle, "device", "")),
)
cached = _SMALLM_TRITON_WEIGHT_WORKSPACES.get(key)
if cached is not None:
return cached
b_shuffle_u8 = b_shuffle.view(torch.uint8).reshape(n // 16, b_shuffle.shape[1] * 16)
b_scale_asm_u8 = b_scale_sh.view(torch.uint8)
scale_rows, scale_cols = b_scale_asm_u8.shape
raw_scales = (
b_scale_asm_u8.view(scale_rows // 32, scale_cols // 8, 4, 16, 2, 2, 1)
.permute(0, 5, 3, 1, 4, 2, 6)
.contiguous()
.view(scale_rows, scale_cols)
)[:n]
b_scale_triton_u8 = (
raw_scales.view(n // 32, 2, 16, scale_cols // 8, 2, 4, 1)
.permute(0, 3, 5, 2, 4, 1, 6)
.contiguous()
.view(n // 32, scale_cols * 32)
)
cached = (b_shuffle_u8, b_scale_triton_u8)
_SMALLM_TRITON_WEIGHT_WORKSPACES[key] = cached
return cached
def _maybe_run_smallm_triton(torch, dtypes, a, b_shuffle, b_scale_sh, m: int, n: int, k: int, plan):
global _SMALLM_TRITON_ERROR, _SMALLM_TRITON_INFO_PRINTED
if not _smallm_triton_enabled(torch, m, n, k, plan):
return None
try:
y = _run_smallm_triton_preshuffle(
torch,
dtypes,
a,
b_shuffle,
b_scale_sh,
m,
n,
)
except Exception:
if _SMALLM_TRITON_ERROR is None:
_SMALLM_TRITON_ERROR = True
try:
import traceback
print("[mxfp4 smallm triton] falling back after error:")
traceback.print_exc()
except Exception:
pass
return None
if not _SMALLM_TRITON_INFO_PRINTED:
_SMALLM_TRITON_INFO_PRINTED = True
try:
print("[mxfp4 smallm triton] using uint8-view preshuffle kernel")
except Exception:
pass
return y
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_raw,
torch::Tensor bs_shuffled,
int64_t m,
int64_t n,
int64_t n_pad);
void mxfp4_quant_shuffle_small_noraw(
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_raw,
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[64];
__shared__ uint8_t fp4_codes[64];
__shared__ uint8_t bs_e8m0[2];
int lane = threadIdx.x;
int row = blockIdx.y;
int subgroup = lane / 32;
int lane_in_group = lane % 32;
int group_base = subgroup * 32;
int block_n = blockIdx.x * 2 + subgroup;
int col = block_n * 32 + lane_in_group;
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_in_group < offset) {
abs_vals[group_base + lane_in_group] =
fmaxf(abs_vals[group_base + lane_in_group], abs_vals[group_base + lane_in_group + offset]);
}
__syncthreads();
}
if (lane_in_group == 0) {
float amax = abs_vals[group_base];
if (amax == 0.0f) {
bs_e8m0[subgroup] = 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[subgroup] = float_to_e8m0(rounded_amax / 4.0f);
}
}
__syncthreads();
float scale = e8m0_to_float(bs_e8m0[subgroup]);
float qx = x_val / scale;
fp4_codes[lane] = float_to_mxfp4(qx);
__syncthreads();
if (lane_in_group < 16) {
uint8_t even = fp4_codes[group_base + lane_in_group * 2];
uint8_t odd = fp4_codes[group_base + lane_in_group * 2 + 1];
x_fp4[row * stride_x_fp4_m + (block_n * 16 + lane_in_group) * stride_x_fp4_n] =
static_cast<uint8_t>(even | (odd << 4));
}
if (lane_in_group == 0) {
if (bs_raw != nullptr) {
bs_raw[row * (n / 32) + block_n] = bs_e8m0[subgroup];
}
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 = 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 + g0 * (32 * n_pad);
bs_shuffled[dst_off] = bs_e8m0[subgroup];
}
}
__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_raw,
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_raw.scalar_type() == torch::kUInt8, "bs_raw must be uint8");
TORCH_CHECK(bs_shuffled.scalar_type() == torch::kUInt8, "bs_shuffled must be uint8");
TORCH_CHECK(
m == 4 || m == 16 || m == 32 || m == 64 || m == 256,
"native quant probe only supports m=4, m=16, m=32, m=64, or m=256");
TORCH_CHECK((n % 64) == 0, "n must be divisible by 64");
const dim3 threads(64);
const dim3 blocks(
static_cast<unsigned int>(n / 64),
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_raw.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));
}
void mxfp4_quant_shuffle_small_noraw(
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 || m == 64 || m == 256,
"native quant probe only supports m=4, m=16, m=32, m=64, or m=256");
TORCH_CHECK((n % 64) == 0, "n must be divisible by 64");
const dim3 threads(64);
const dim3 blocks(
static_cast<unsigned int>(n / 64),
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>(),
nullptr,
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_v4",
cpp_sources=[cpp_src],
cuda_sources=[hip_src],
functions=[
"mxfp4_quant_shuffle_small",
"mxfp4_quant_shuffle_small_noraw",
"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_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_noraw(
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_with_workspace(torch, dtypes, x, workspace):
triton, quant_kernel, shuffle_kernel = _get_specialized_quant_runtime()
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_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 _quant_mxfp4_specialized(torch, dtypes, x):
workspace = _get_specialized_quant_workspace(torch, x)
return _quant_mxfp4_specialized_with_workspace(torch, dtypes, x, workspace)
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)
smallm_triton = _maybe_run_smallm_triton(
torch,
dtypes,
a,
b_shuffle,
b_scale_sh,
m,
n,
k,
plan,
)
if smallm_triton is not None:
return smallm_triton
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 · 1097 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