submission 693471
nikxkilla · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 672 lines, June 9 Researcher Reciprocity License v1.0.
submission_v3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-693471?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:01f8afa7c3322fa771342e686691a4512c5622d25a7da4976451f574fb54bcef
license declaredunknown
license concludedunknown
authorsnikxkilla
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
m.def("quant_into_shuffled_hip", &quant_into_shuffled_hip, "HIP small-m native fp4 quant");shared-memory
__shared__ float s_vals[GROUPS_PER_BLOCK][32];split-k
_CUSTOM_CSV = """cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatiostages = 1
num_stages=1,Kernel source
submission_v3.py672 lines
from task import input_t, output_t
import functools
import os
import tempfile
import textwrap
import triton
import triton.language as tl
if "PYTORCH_ROCM_ARCH" not in os.environ:
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
_WS = {}
_CSV_PATH = None
_CUSTOM_CSV = """cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio
256,4,2880,512,21,0,0.0,_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E,0.0,0.0,0.0
256,16,2112,7168,21,0,0.0,_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E,0.0,0.0,0.0
256,32,4096,512,21,0,0.0,_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E,0.0,0.0,0.0
256,32,2880,512,21,0,0.0,_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E,0.0,0.0,0.0
256,64,7168,2048,21,0,6.8112,_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E,275.88,1221.97,0.0
256,256,3072,1536,21,0,6.1771,_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E,391.11,668.4,0.0
"""
_HIP_SRC = textwrap.dedent(
r"""
#include <torch/extension.h>
#include <pybind11/pybind11.h>
#include <hip/hip_runtime.h>
#include <hip/amd_detail/amd_hip_bf16.h>
#include <cstdint>
#include <stdexcept>
#define CHECK_CUDA(x) TORCH_CHECK(x.is_cuda(), #x " must be on ROCm device")
#define CHECK_CONTIGUOUS(x) TORCH_CHECK(x.is_contiguous(), #x " must be contiguous")
#define CHECK_BF16(x) TORCH_CHECK(x.scalar_type() == torch::kBFloat16, #x " must be bf16")
#define CHECK_U8(x) TORCH_CHECK(x.scalar_type() == torch::kUInt8, #x " must be uint8")
static __device__ inline uint32_t f32_as_u32(float x) {
return __builtin_bit_cast(uint32_t, x);
}
static __device__ inline int shuffled_scale_index(
int row_idx, int scale_n_idx, int scale_n_pad) {
int a = row_idx / 32;
int m1 = row_idx % 32;
int b = m1 / 16;
int c = m1 % 16;
int d = scale_n_idx / 8;
int n1 = scale_n_idx % 8;
int e = n1 / 4;
int f = n1 % 4;
int gn = scale_n_pad / 8;
return (((((a * gn + d) * 4 + f) * 16 + c) * 2 + e) * 2 + b);
}
static __device__ inline uint8_t quant_one_fp4_e2m1(float x, uint8_t scale_e8m0) {
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;
int scale_unbiased = static_cast<int>(scale_e8m0) - 127;
float qx = x * exp2f(-static_cast<float>(scale_unbiased));
uint32_t qx_bits = f32_as_u32(qx);
uint32_t sign = qx_bits & 0x80000000u;
uint32_t abs_bits = qx_bits ^ sign;
float abs_f = __builtin_bit_cast(float, abs_bits);
bool saturate_mask = abs_f >= MAX_NORMAL;
bool denormal_mask = (!saturate_mask) && (abs_f < MIN_NORMAL);
uint8_t fp4_val = 0x7u;
if (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 = __builtin_bit_cast(float, denorm_mask_int);
uint32_t denormal_x = f32_as_u32(abs_f + denorm_mask_float);
denormal_x -= denorm_mask_int;
fp4_val = static_cast<uint8_t>(denormal_x);
} else if (!saturate_mask) {
uint32_t mant_odd = (abs_bits >> (MBITS_F32 - MBITS_FP4)) & 1u;
constexpr int32_t val_to_add =
((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1 << 21) - 1;
int32_t normal_x = static_cast<int32_t>(abs_bits);
normal_x += val_to_add;
normal_x += static_cast<int32_t>(mant_odd);
normal_x >>= (MBITS_F32 - MBITS_FP4);
fp4_val = static_cast<uint8_t>(normal_x);
}
uint8_t sign_lp = static_cast<uint8_t>(
sign >> (MBITS_F32 + EBITS_F32 - MBITS_FP4 - EBITS_FP4));
return static_cast<uint8_t>(fp4_val | sign_lp);
}
static __device__ inline uint8_t encode_scale_from_amax(float amax) {
if (!(amax > 0.0f)) {
return static_cast<uint8_t>(0);
}
uint32_t amax_bits = f32_as_u32(amax);
uint32_t rounded = (amax_bits + 0x200000u) & 0xFF800000u;
float rounded_amax = __builtin_bit_cast(float, rounded);
int scale_unbiased = static_cast<int>(floorf(log2f(rounded_amax)) - 2.0f);
if (scale_unbiased < -127) scale_unbiased = -127;
if (scale_unbiased > 127) scale_unbiased = 127;
return static_cast<uint8_t>(scale_unbiased + 127);
}
template <int GROUPS_PER_BLOCK>
__global__ __launch_bounds__(GROUPS_PER_BLOCK * 32)
void quant_smallm_native_fp4_kernel(
const __hip_bfloat16* __restrict__ x_ptr,
uint8_t* __restrict__ x_fp4_ptr,
uint8_t* __restrict__ bs_shuffled_ptr,
int64_t stride_x_m,
int64_t stride_x_n,
int64_t stride_x_fp4_m,
int64_t stride_x_fp4_n,
int M,
int K,
int SCALE_N_VALID,
int SCALE_N_PAD) {
__shared__ float s_vals[GROUPS_PER_BLOCK][32];
__shared__ float s_abs[GROUPS_PER_BLOCK][32];
__shared__ uint8_t s_codes[GROUPS_PER_BLOCK][32];
__shared__ uint8_t s_scale[GROUPS_PER_BLOCK];
const int tid = threadIdx.x;
const int grp = tid >> 5;
const int lane = tid & 31;
const int row = blockIdx.x * GROUPS_PER_BLOCK + grp;
const int scale_n = blockIdx.y;
const int k_idx = scale_n * 32 + lane;
float v = 0.0f;
if (grp < GROUPS_PER_BLOCK && row < M && scale_n < SCALE_N_VALID && k_idx < K) {
v = static_cast<float>(x_ptr[row * stride_x_m + k_idx * stride_x_n]);
}
s_vals[grp][lane] = v;
s_abs[grp][lane] = fabsf(v);
__syncthreads();
for (int offset = 16; offset > 0; offset >>= 1) {
if (lane < offset) {
float other = s_abs[grp][lane + offset];
if (other > s_abs[grp][lane]) {
s_abs[grp][lane] = other;
}
}
__syncthreads();
}
if (lane == 0) {
uint8_t scale_e8m0 = encode_scale_from_amax(s_abs[grp][0]);
s_scale[grp] = scale_e8m0;
if (row < M && scale_n < SCALE_N_VALID) {
int lin = shuffled_scale_index(row, scale_n, SCALE_N_PAD);
bs_shuffled_ptr[lin] = scale_e8m0;
}
}
__syncthreads();
if (row < M && scale_n < SCALE_N_VALID) {
s_codes[grp][lane] = quant_one_fp4_e2m1(s_vals[grp][lane], s_scale[grp]) & 0x0Fu;
} else {
s_codes[grp][lane] = 0;
}
__syncthreads();
if ((lane & 1) == 0 && row < M && scale_n < SCALE_N_VALID) {
uint8_t packed = static_cast<uint8_t>(
s_codes[grp][lane] | (s_codes[grp][lane + 1] << 4));
int out_n = scale_n * 16 + (lane >> 1);
x_fp4_ptr[row * stride_x_fp4_m + out_n * stride_x_fp4_n] = packed;
}
}
void quant_into_shuffled_hip(
torch::Tensor A,
torch::Tensor A_q_u8,
torch::Tensor A_scale_sh_u8) {
CHECK_CUDA(A);
CHECK_CUDA(A_q_u8);
CHECK_CUDA(A_scale_sh_u8);
CHECK_CONTIGUOUS(A);
CHECK_CONTIGUOUS(A_q_u8);
CHECK_CONTIGUOUS(A_scale_sh_u8);
CHECK_BF16(A);
CHECK_U8(A_q_u8);
CHECK_U8(A_scale_sh_u8);
TORCH_CHECK(A.dim() == 2, "A must be 2D");
TORCH_CHECK(A_q_u8.dim() == 2, "A_q_u8 must be 2D");
TORCH_CHECK(A_scale_sh_u8.dim() == 2, "A_scale_sh_u8 must be 2D");
const int M = static_cast<int>(A.size(0));
const int K = static_cast<int>(A.size(1));
const int SCALE_N_VALID = (K + 31) / 32;
const int SCALE_N_PAD = static_cast<int>(A_scale_sh_u8.size(1));
TORCH_CHECK(M <= 32, "HIP small-m quant path only supports m <= 32");
TORCH_CHECK(A_q_u8.size(0) == M, "A_q_u8 row mismatch");
TORCH_CHECK(A_q_u8.size(1) * 2 >= K, "A_q_u8 packed width mismatch");
constexpr int GROUPS_PER_BLOCK = 4;
dim3 block(GROUPS_PER_BLOCK * 32);
dim3 grid((M + GROUPS_PER_BLOCK - 1) / GROUPS_PER_BLOCK, SCALE_N_VALID);
hipLaunchKernelGGL(
(quant_smallm_native_fp4_kernel<GROUPS_PER_BLOCK>),
grid,
block,
0,
0,
reinterpret_cast<const __hip_bfloat16*>(A.data_ptr<at::BFloat16>()),
reinterpret_cast<uint8_t*>(A_q_u8.data_ptr<uint8_t>()),
reinterpret_cast<uint8_t*>(A_scale_sh_u8.data_ptr<uint8_t>()),
static_cast<int64_t>(A.stride(0)),
static_cast<int64_t>(A.stride(1)),
static_cast<int64_t>(A_q_u8.stride(0)),
static_cast<int64_t>(A_q_u8.stride(1)),
M,
K,
SCALE_N_VALID,
SCALE_N_PAD);
auto err = hipGetLastError();
if (err != hipSuccess) {
throw std::runtime_error(hipGetErrorString(err));
}
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("quant_into_shuffled_hip", &quant_into_shuffled_hip, "HIP small-m native fp4 quant");
}
"""
)
@functools.lru_cache(maxsize=1)
def _load_hip_smallm_mod():
from torch.utils.cpp_extension import load_inline
return load_inline(
name="mxfp4_quant_smallm_nativefp4_gfx950",
cpp_sources="",
cuda_sources=_HIP_SRC,
functions=None,
with_cuda=True,
verbose=False,
extra_cuda_cflags=["-O3", "-std=c++20", "--offload-arch=gfx950"],
no_implicit_headers=False,
)
def _ensure_csv():
global _CSV_PATH
if _CSV_PATH is None:
f = tempfile.NamedTemporaryFile(
mode="w",
suffix=".csv",
delete=False,
prefix="aiter_a4w4_fused_scale_codex_",
)
f.write(_CUSTOM_CSV)
f.close()
_CSV_PATH = f.name
os.environ["AITER_CONFIG_GEMM_A4W4"] = _CSV_PATH
@triton.jit
def _mxfp4_quant_op_local(
x,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_M: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
):
EXP_BIAS_FP32: tl.constexpr = 127
EXP_BIAS_FP4: tl.constexpr = 1
EBITS_F32: tl.constexpr = 8
EBITS_FP4: tl.constexpr = 2
MBITS_F32: tl.constexpr = 23
MBITS_FP4: tl.constexpr = 1
max_normal: tl.constexpr = 6
min_normal: tl.constexpr = 1
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
amax = amax.to(tl.int32, bitcast=True)
amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax = amax.to(tl.float32, bitcast=True)
scale_e8m0_unbiased = tl.log2(amax).floor() - 2
scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
quant_scale = tl.exp2(-scale_e8m0_unbiased)
qx = x * quant_scale
qx = qx.to(tl.uint32, bitcast=True)
s = qx & 0x80000000
qx = qx ^ s
qx_fp32 = qx.to(tl.float32, bitcast=True)
saturate_mask = qx_fp32 >= max_normal
denormal_mask = (~saturate_mask) & (qx_fp32 < min_normal)
normal_mask = ~(saturate_mask | denormal_mask)
denorm_exp: tl.constexpr = (
(EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1
)
denorm_mask_int: tl.constexpr = denorm_exp << MBITS_F32
denorm_mask_float: tl.constexpr = tl.cast(denorm_mask_int, tl.float32, bitcast=True)
denormal_x = qx_fp32 + denorm_mask_float
denormal_x = denormal_x.to(tl.uint32, bitcast=True)
denormal_x -= denorm_mask_int
denormal_x = denormal_x.to(tl.uint8)
normal_x = qx
mant_odd = (normal_x >> (MBITS_F32 - MBITS_FP4)) & 1
val_to_add = ((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1 << 21) - 1
normal_x += val_to_add
normal_x += mant_odd
normal_x = normal_x >> (MBITS_F32 - MBITS_FP4)
normal_x = normal_x.to(tl.uint8)
e2m1_value = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)
e2m1_value = tl.where(normal_mask, normal_x, e2m1_value)
e2m1_value = tl.where(denormal_mask, denormal_x, e2m1_value)
sign_lp = s >> (MBITS_F32 + EBITS_F32 - MBITS_FP4 - EBITS_FP4)
sign_lp = sign_lp.to(tl.uint8)
e2m1_value = e2m1_value | sign_lp
e2m1_value = tl.reshape(
e2m1_value, [BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2]
)
evens, odds = tl.split(e2m1_value)
x_fp4 = evens | (odds << 4)
x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
@triton.heuristics(
{
"EVEN_M_N": lambda args: args["M"] % args["BLOCK_SIZE_M"] == 0
and args["N"] % (args["BLOCK_SIZE_N"] * args["NUM_ITER"]) == 0,
}
)
@triton.jit
def _dynamic_mxfp4_quant_kernel_shuffled(
x_ptr,
x_fp4_ptr,
bs_shuffled_ptr,
stride_x_m_in,
stride_x_n_in,
stride_x_fp4_m_in,
stride_x_fp4_n_in,
M,
N,
SCALE_N_VALID,
SCALE_N_PAD,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
NUM_ITER: tl.constexpr,
NUM_STAGES: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
EVEN_M_N: tl.constexpr,
):
pid_m = tl.program_id(0)
start_n = tl.program_id(1) * NUM_ITER
stride_x_m = tl.cast(stride_x_m_in, tl.int64)
stride_x_n = tl.cast(stride_x_n_in, tl.int64)
stride_x_fp4_m = tl.cast(stride_x_fp4_m_in, tl.int64)
stride_x_fp4_n = tl.cast(stride_x_fp4_n_in, tl.int64)
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
for pid_n in tl.range(start_n, start_n + NUM_ITER, num_stages=NUM_STAGES):
x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
if EVEN_M_N:
x = tl.load(x_ptr + x_offs, cache_modifier=".cg").to(tl.float32)
else:
x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
x = tl.load(x_ptr + x_offs, mask=x_mask, cache_modifier=".cg").to(tl.float32)
out_tensor, bs_e8m0 = _mxfp4_quant_op_local(
x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE
)
out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
out_offs_n = pid_n * (BLOCK_SIZE_N // 2) + tl.arange(0, BLOCK_SIZE_N // 2)
out_offs = out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
if EVEN_M_N:
tl.store(x_fp4_ptr + out_offs, out_tensor)
else:
out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)
bs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
bs_n = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
a = bs_m[:, None] // 32
m1 = bs_m[:, None] % 32
b = m1 // 16
c = m1 % 16
d = bs_n[None, :] // 8
n1 = bs_n[None, :] % 8
e = n1 // 4
f = n1 % 4
gn = SCALE_N_PAD // 8
lin = (((((a * gn + d) * 4 + f) * 16 + c) * 2 + e) * 2 + b)
bs_mask = (bs_m[:, None] < M) & (bs_n[None, :] < SCALE_N_VALID)
tl.store(bs_shuffled_ptr + tl.cast(lin, tl.int64), bs_e8m0, mask=bs_mask)
def _get_ws(device, m, n, k):
import torch
from aiter import dtypes
key = (device.index if device.index is not None else 0, m, n, k)
ws = _WS.get(key)
if ws is None:
scale_n = (k + 31) // 32
scale_n_pad = ((scale_n + 7) // 8) * 8
scale_m_pad = ((m + 255) // 256) * 256
ascale_cm_buf_u8 = torch.empty((scale_n, m), dtype=torch.uint8, device=device)
ascale_sh_u8 = torch.empty(
(scale_m_pad, scale_n_pad), dtype=torch.uint8, device=device
)
ascale_sh_u8.fill_(127)
ws = {
"aq_u8": torch.empty((m, k // 2), dtype=torch.uint8, device=device),
"ascale_cm_buf_u8": ascale_cm_buf_u8,
"ascale_cm_u8": ascale_cm_buf_u8.T,
"ascale_sh_u8": ascale_sh_u8,
"out": torch.empty(
(((m + 31) // 32) * 32, n), dtype=dtypes.bf16, device=device
),
}
_WS[key] = ws
return ws
ACTIVE_SWEEP = "s7"
_QUANT_SWEEPS = {
"s0": {
"small": {"NUM_ITER": 1, "BLOCK_SIZE_M": "pow2_cap32", "BLOCK_SIZE_N": 256, "NUM_WARPS": 4, "NUM_STAGES": 1},
"mid": {"NUM_ITER": 4, "BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 2},
"large": {"NUM_ITER": 4, "BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 2},
},
"s1": {
"small": {"NUM_ITER": 1, "BLOCK_SIZE_M": "pow2_cap32", "BLOCK_SIZE_N": 256, "NUM_WARPS": 2, "NUM_STAGES": 1},
"mid": {"NUM_ITER": 4, "BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 2},
"large": {"NUM_ITER": 4, "BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 2},
},
"s2": {
"small": {"NUM_ITER": 1, "BLOCK_SIZE_M": "pow2_cap32", "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 1},
"mid": {"NUM_ITER": 4, "BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 2},
"large": {"NUM_ITER": 4, "BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 2},
},
"s3": {
"small": {"NUM_ITER": 1, "BLOCK_SIZE_M": "pow2_cap32", "BLOCK_SIZE_N": 256, "NUM_WARPS": 4, "NUM_STAGES": 1},
"mid": {"NUM_ITER": 2, "BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 2},
"large": {"NUM_ITER": 4, "BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 2},
},
"s4": {
"small": {"NUM_ITER": 1, "BLOCK_SIZE_M": "pow2_cap32", "BLOCK_SIZE_N": 256, "NUM_WARPS": 4, "NUM_STAGES": 1},
"mid": {"NUM_ITER": 4, "BLOCK_SIZE_M": 64, "BLOCK_SIZE_N": 64, "NUM_WARPS": 4, "NUM_STAGES": 2},
"large": {"NUM_ITER": 4, "BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 2},
},
"s5": {
"small": {"NUM_ITER": 1, "BLOCK_SIZE_M": "pow2_cap32", "BLOCK_SIZE_N": 256, "NUM_WARPS": 4, "NUM_STAGES": 1},
"mid": {"NUM_ITER": 4, "BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 2},
"large": {"NUM_ITER": 4, "BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 64, "NUM_WARPS": 4, "NUM_STAGES": 2},
},
"s6": {
"small": {"NUM_ITER": 1, "BLOCK_SIZE_M": "pow2_cap32", "BLOCK_SIZE_N": 256, "NUM_WARPS": 4, "NUM_STAGES": 1},
"mid": {"NUM_ITER": 4, "BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 2},
"large": {"NUM_ITER": 4, "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 64, "NUM_WARPS": 4, "NUM_STAGES": 2},
},
"s7": {
"small": {"NUM_ITER": 1, "BLOCK_SIZE_M": "pow2_cap32", "BLOCK_SIZE_N": 256, "NUM_WARPS": 4, "NUM_STAGES": 1},
"mid": {"NUM_ITER": 4, "BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "NUM_WARPS": 4, "NUM_STAGES": 2},
"large": {"NUM_ITER": 4, "BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "NUM_WARPS": 8, "NUM_STAGES": 2},
},
}
ACTIVE_LARGE = "L1"
_LARGE_K_CFGS = {
"L0": {
"NUM_ITER": 4,
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"NUM_WARPS": 4,
"NUM_STAGES": 2,
},
"L1": {
"NUM_ITER": 2,
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"NUM_WARPS": 4,
"NUM_STAGES": 2,
},
"L2": {
"NUM_ITER": 4,
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 64,
"NUM_WARPS": 4,
"NUM_STAGES": 2,
},
}
def _resolve_block_m(m, block_m):
import triton
if block_m == "pow2":
return triton.next_power_of_2(m)
if block_m == "pow2_cap32":
return min(32, triton.next_power_of_2(m))
return block_m
def _pick_quant_cfg(M, K):
if K <= 1024:
src = _QUANT_SWEEPS["s2"]["small"]
elif K <= 4096:
src = _QUANT_SWEEPS["s3"]["mid"]
else:
src = _LARGE_K_CFGS[ACTIVE_LARGE]
cfg = dict(src)
cfg["BLOCK_SIZE_M"] = _resolve_block_m(M, cfg["BLOCK_SIZE_M"])
return cfg
def _quant_into_raw(A, aq_u8, ascale_u8):
from aiter.ops.triton._triton_kernels.quant.quant import _dynamic_mxfp4_quant_kernel
M, K = A.shape
cfg = _pick_quant_cfg(M, K)
grid = (
triton.cdiv(M, cfg["BLOCK_SIZE_M"]),
triton.cdiv(K, cfg["BLOCK_SIZE_N"] * cfg["NUM_ITER"]),
)
_dynamic_mxfp4_quant_kernel[grid](
A,
aq_u8,
ascale_u8,
*A.stride(),
*aq_u8.stride(),
*ascale_u8.stride(),
M=M,
N=K,
MXFP4_QUANT_BLOCK_SIZE=32,
SCALING_MODE=0,
NUM_ITER=cfg["NUM_ITER"],
BLOCK_SIZE_M=cfg["BLOCK_SIZE_M"],
BLOCK_SIZE_N=cfg["BLOCK_SIZE_N"],
NUM_STAGES=cfg["NUM_STAGES"],
num_warps=cfg["NUM_WARPS"],
waves_per_eu=0,
num_stages=1,
)
return aq_u8, ascale_u8
def _quant_into_shuffled_triton(A, aq_u8, ascale_sh_u8):
M, K = A.shape
cfg = _pick_quant_cfg(M, K)
scale_n_valid = (K + 31) // 32
scale_n_pad = ascale_sh_u8.shape[1]
grid = (
triton.cdiv(M, cfg["BLOCK_SIZE_M"]),
triton.cdiv(K, cfg["BLOCK_SIZE_N"] * cfg["NUM_ITER"]),
)
_dynamic_mxfp4_quant_kernel_shuffled[grid](
A,
aq_u8,
ascale_sh_u8,
*A.stride(),
*aq_u8.stride(),
M=M,
N=K,
SCALE_N_VALID=scale_n_valid,
SCALE_N_PAD=scale_n_pad,
MXFP4_QUANT_BLOCK_SIZE=32,
NUM_ITER=cfg["NUM_ITER"],
BLOCK_SIZE_M=cfg["BLOCK_SIZE_M"],
BLOCK_SIZE_N=cfg["BLOCK_SIZE_N"],
NUM_STAGES=cfg["NUM_STAGES"],
num_warps=cfg["NUM_WARPS"],
waves_per_eu=0,
num_stages=1,
)
return aq_u8, ascale_sh_u8
def _quant_into_shuffled_smallm_hip(A, aq_u8, ascale_sh_u8):
mod = _load_hip_smallm_mod()
mod.quant_into_shuffled_hip(A, aq_u8, ascale_sh_u8)
return aq_u8, ascale_sh_u8
def _quant_into_shuffled(A, aq_u8, ascale_sh_u8):
if A.shape[0] <= 32:
return _quant_into_shuffled_smallm_hip(A, aq_u8, ascale_sh_u8)
return _quant_into_shuffled_triton(A, aq_u8, ascale_sh_u8)
def custom_kernel(data: input_t) -> output_t:
_ensure_csv()
import aiter
from aiter import dtypes
A, _, _, B_shuffle, B_scale_sh = data
A = A.contiguous()
B_shuffle = B_shuffle.contiguous()
B_scale_sh = B_scale_sh.contiguous()
m, k = A.shape
n = B_shuffle.shape[0]
ws = _get_ws(A.device, m, n, k)
A_q_u8, A_scale_sh_u8 = _quant_into_shuffled(A, ws["aq_u8"], ws["ascale_sh_u8"])
A_scale = A_scale_sh_u8.view(dtypes.fp8_e8m0)
A_q = A_q_u8.view(dtypes.fp4x2)
return aiter.gemm_a4w4(
A_q,
B_shuffle,
A_scale,
B_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
scrolls · 672 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