submission 564882
trungnob · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1082 lines, June 9 Researcher Reciprocity License v1.0.
submission_v_23.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-564882?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:ec3af872bf871deb44a1e888b12ecfdc49a9adecd626f308e402e11e6726bd37
license declaredunknown
license concludedunknown
authorstrungnob
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MXFP4 GEMM hybrid:shared-memory
__shared__ float s_abs[WAVE_SIZE];split-k
"SPLITK_BLOCK_SIZE": 2 * 7168 // 4,Kernel source
submission_v_23.py1082 lines
"""
MXFP4 GEMM hybrid:
- monkey-patched aiter.gemm_a16wfp4 for shapes where direct BF16-A/FP4-B wins
- HIP wavefront quant + tuned ASM GEMM where the direct path loses
"""
import os
from typing import Optional
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
import torch
import triton
import triton.language as tl
from torch.utils.cpp_extension import load_inline
import aiter
import aiter.ops.triton.utils._triton.arch_info as arch_info
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm, get_padded_m
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
_gemm_afp4wfp4_reduce_kernel,
)
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
from aiter.ops.triton.utils._triton.kernel_repr import make_kernel_repr
from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid
from aiter.ops.triton.utils.common_utils import deserialize_str
from aiter.ops.triton.utils.gemm_config_utils import get_gemm_config
from task import input_t, output_t
TUNED_CONFIGS = {
(4, 2880, 512): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x256E", 3),
(8, 2112, 7168): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E", 2),
(16, 2112, 7168): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E", 1),
(16, 3072, 1536): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 0),
(32, 4096, 512): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_96x384E", 1),
(32, 2880, 512): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 0),
(64, 7168, 2048): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E", 1),
(64, 3072, 1536): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 0),
(256, 2880, 512): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x384E", 2),
(256, 3072, 1536): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E", 1),
}
DIRECT_SHAPES = {
(4, 2880, 512),
(32, 4096, 512),
(32, 2880, 512),
(64, 7168, 2048),
}
OUTLIER_RAW_SHAPE = (16, 2112, 7168)
OUTLIER_RAW_CONFIG = {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
# v7 candidate: keep the fused outlier route but halve the K-loop count.
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 1,
"NUM_KSPLIT": 4,
"SPLITK_BLOCK_SIZE": 2 * 7168 // 4,
"num_warps": 8,
"num_stages": 2,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
}
HIP_SRC = r"""
#include <torch/extension.h>
#include <hip/hip_runtime.h>
#include <hip/hip_bf16.h>
#include <stdint.h>
#define SCALE_GROUP_SIZE 32
#define WAVE_SIZE 256
#define GROUPS_PER_BLOCK (WAVE_SIZE / 32)
__device__ __forceinline__ uint8_t fp32_to_e2m1_exact(float qv) {
constexpr uint32_t EXP_BIAS_FP32 = 127;
constexpr uint32_t EXP_BIAS_FP4 = 1;
constexpr float max_normal = 6.0f;
constexpr float min_normal = 1.0f;
uint32_t bits = __float_as_uint(qv);
uint32_t sign = bits & 0x80000000u;
bits ^= sign;
float absv = __uint_as_float(bits);
uint8_t result;
if (absv >= max_normal) {
result = 0x7;
} else if (absv < min_normal) {
constexpr uint32_t denorm_exp =
(EXP_BIAS_FP32 - EXP_BIAS_FP4) + (23 - 1) + 1;
constexpr uint32_t denorm_mask_int = denorm_exp << 23;
float denorm_mask_float = __uint_as_float(denorm_mask_int);
float denormal_x = absv + denorm_mask_float;
uint32_t denormal_bits = __float_as_uint(denormal_x);
denormal_bits -= denorm_mask_int;
result = static_cast<uint8_t>(denormal_bits & 0xFF);
} else {
uint32_t normal_x = bits;
uint32_t mant_odd = (normal_x >> (23 - 1)) & 1;
int32_t val_to_add =
((static_cast<int32_t>(EXP_BIAS_FP4 - EXP_BIAS_FP32)) << 23) + (1 << 21) - 1;
normal_x += static_cast<uint32_t>(val_to_add);
normal_x += mant_odd;
normal_x >>= (23 - 1);
result = static_cast<uint8_t>(normal_x & 0xFF);
}
uint8_t sign_lp = static_cast<uint8_t>(
(sign >> (23 + 8 - 1 - 2)) & 0x8
);
return result | sign_lp;
}
__device__ __forceinline__ uint8_t compute_e8m0_scale_exact(float amax) {
if (amax == 0.0f) {
return 127;
}
uint32_t bits = __float_as_uint(amax);
bits = (bits + 0x200000u) & 0xFF800000u;
int scale_unbiased = static_cast<int>((bits >> 23) & 0xFF) - 127 - 2;
if (scale_unbiased < -127) scale_unbiased = -127;
if (scale_unbiased > 127) scale_unbiased = 127;
return static_cast<uint8_t>(scale_unbiased + 127);
}
__global__ __launch_bounds__(WAVE_SIZE)
void fused_quant_wave_kernel(
const __hip_bfloat16* __restrict__ A,
uint8_t* __restrict__ A_fp4,
uint8_t* __restrict__ A_scale,
int M,
int K,
int scaleN,
int scaleN_pad
) {
__shared__ float s_abs[WAVE_SIZE];
__shared__ uint8_t s_nib[WAVE_SIZE];
int row = blockIdx.x;
int group_block = blockIdx.y;
int lane = threadIdx.x;
if (row >= M) {
return;
}
int half = lane >> 5;
int lane32 = lane & 31;
int group = group_block * GROUPS_PER_BLOCK + half;
bool valid_group = group < scaleN;
float v = 0.0f;
if (valid_group) {
int k_idx = group * SCALE_GROUP_SIZE + lane32;
v = __bfloat162float(A[row * K + k_idx]);
}
s_abs[lane] = fabsf(v);
__syncthreads();
int half_base = half * 32;
for (int offset = 16; offset > 0; offset >>= 1) {
if (lane32 < offset) {
float other = s_abs[half_base + lane32 + offset];
s_abs[half_base + lane32] = fmaxf(s_abs[half_base + lane32], other);
}
__syncthreads();
}
uint8_t scale_e8m0 = compute_e8m0_scale_exact(s_abs[half_base]);
float scale_unbiased = static_cast<float>(static_cast<int>(scale_e8m0) - 127);
float inv_scale = exp2f(-scale_unbiased);
s_nib[lane] = fp32_to_e2m1_exact(v * inv_scale);
__syncthreads();
if (valid_group && lane32 < 16) {
int nib_base = half_base + lane32 * 2;
uint8_t lo = s_nib[nib_base];
uint8_t hi = s_nib[nib_base + 1];
A_fp4[row * (K / 2) + group * 16 + lane32] = lo | (hi << 4);
}
if (valid_group && lane32 == 0) {
int d0 = row / 32;
int rem_r = row % 32;
int d1 = rem_r / 16;
int d2 = rem_r % 16;
int d3 = group / 8;
int rem_c = group % 8;
int d4 = rem_c / 4;
int d5 = rem_c % 4;
int sn8 = scaleN_pad / 8;
int shuffle_idx = d0 * (sn8 * 256) + d3 * 256 + d5 * 64 + d2 * 4 + d4 * 2 + d1;
A_scale[shuffle_idx] = scale_e8m0;
}
}
void fused_mxfp4_quant_out(
torch::Tensor A,
torch::Tensor A_fp4,
torch::Tensor A_scale,
int scaleN,
int scaleN_pad
) {
int M = static_cast<int>(A.size(0));
int K = static_cast<int>(A.size(1));
dim3 blocks(M, (scaleN + GROUPS_PER_BLOCK - 1) / GROUPS_PER_BLOCK);
dim3 threads(WAVE_SIZE);
fused_quant_wave_kernel<<<blocks, threads>>>(
reinterpret_cast<const __hip_bfloat16*>(A.data_ptr()),
A_fp4.data_ptr<uint8_t>(),
A_scale.data_ptr<uint8_t>(),
M,
K,
scaleN,
scaleN_pad
);
}
"""
CPP_SRC = r"""
#include <torch/extension.h>
void fused_mxfp4_quant_out(
torch::Tensor A,
torch::Tensor A_fp4,
torch::Tensor A_scale,
int scaleN,
int scaleN_pad
);
"""
_module = None
_hip_buf_cache = {}
_index_cache = {}
_raw_scale_buf_cache = {}
_view_scale_cache = {}
_direct_out_cache = {}
_fused_pp_cache = {}
_patched = False
_status_printed = set()
def _status(key, message: str) -> None:
if key not in _status_printed:
_status_printed.add(key)
print(f"[CODEX-HYBRID] {message}", flush=True)
@triton.jit
def _unshuffle_e8m0_kernel(
scale_sh_ptr,
raw_ptr,
sm: tl.constexpr,
sn: tl.constexpr,
rows: tl.constexpr,
cols: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
):
pid = tl.program_id(0)
offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
mask = offsets < (rows * cols)
sn8: tl.constexpr = sn // 8
r = offsets // cols
c = offsets % cols
d0 = r // 32
rem_r = r % 32
d1 = rem_r // 16
d2 = rem_r % 16
d3 = c // 8
rem_c = c % 8
d4 = rem_c // 4
d5 = rem_c % 4
shuffled_idx = d0 * (sn8 * 256) + d3 * 256 + d5 * 64 + d2 * 4 + d4 * 2 + d1
val = tl.load(scale_sh_ptr + shuffled_idx, mask=mask)
tl.store(raw_ptr + offsets, val, mask=mask)
def fast_unshuffle_e8m0(scale_sh: torch.Tensor, rows: int, cols: int) -> torch.Tensor:
sm, sn = scale_sh.shape
device_index = scale_sh.device.index if scale_sh.device.index is not None else 0
key = (device_index, rows, cols)
raw_flat = _raw_scale_buf_cache.get(key)
if raw_flat is None:
raw_flat = torch.empty(rows * cols, dtype=torch.uint8, device=scale_sh.device)
_raw_scale_buf_cache[key] = raw_flat
block_size = 1024
grid = (triton.cdiv(rows * cols, block_size),)
_unshuffle_e8m0_kernel[grid](
scale_sh.view(torch.uint8),
raw_flat,
sm,
sn,
rows,
cols,
BLOCK_SIZE=block_size,
)
return raw_flat.view(rows, cols).view(scale_sh.dtype)
_gemm_a16wfp4_fused_repr = make_kernel_repr(
"_gemm_a16wfp4_fused_kernel",
[
"BLOCK_SIZE_M",
"BLOCK_SIZE_N",
"BLOCK_SIZE_K",
"GROUP_SIZE_M",
"num_warps",
"num_stages",
"waves_per_eu",
"matrix_instr_nonkdim",
"cache_modifier",
"NUM_KSPLIT",
],
)
@triton.heuristics(
{
"EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)
and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),
"GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])
* triton.cdiv(args["N"], args["BLOCK_SIZE_N"]),
}
)
@triton.jit(repr=_gemm_a16wfp4_fused_repr)
def _gemm_a16wfp4_fused_kernel(
a_ptr,
b_ptr,
c_ptr,
b_scales_ptr,
M,
N,
K,
stride_am,
stride_ak,
stride_bk,
stride_bn,
stride_ck,
stride_cm,
stride_cn,
scale_sh_cols,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr,
NUM_KSPLIT: tl.constexpr,
SPLITK_BLOCK_SIZE: tl.constexpr,
EVEN_K: tl.constexpr,
num_warps: tl.constexpr,
num_stages: tl.constexpr,
waves_per_eu: tl.constexpr,
matrix_instr_nonkdim: tl.constexpr,
GRID_MN: tl.constexpr,
ATOMIC_ADD: tl.constexpr,
cache_modifier: tl.constexpr,
):
tl.assume(stride_am > 0)
tl.assume(stride_ak > 0)
tl.assume(stride_bk > 0)
tl.assume(stride_bn > 0)
tl.assume(stride_cm > 0)
tl.assume(stride_cn > 0)
tl.assume(scale_sh_cols > 0)
pid_unified = tl.program_id(axis=0)
pid_k = pid_unified % NUM_KSPLIT
pid = pid_unified // NUM_KSPLIT
num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
if NUM_KSPLIT == 1:
pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
else:
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
SCALE_GROUP_SIZE: tl.constexpr = 32
if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16
offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
a_ptrs = a_ptr + (
offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak
)
offs_k = tl.arange(0, BLOCK_SIZE_K // 2)
offs_k_split = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_k
offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
b_ptrs = b_ptr + (offs_k_split[:, None] * stride_bk + offs_bn[None, :] * stride_bn)
ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE)) + tl.arange(
0, BLOCK_SIZE_K // SCALE_GROUP_SIZE
)
offs_bn_2d = offs_bn[:, None]
ks_2d = ks[None, :]
sn8 = scale_sh_cols // 8
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for _ in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
shuffle_idx = (
(offs_bn_2d // 32) * (sn8 * 256)
+ (ks_2d // 8) * 256
+ ((ks_2d % 8) % 4) * 64
+ (offs_bn_2d % 16) * 4
+ ((ks_2d % 8) // 4) * 2
+ ((offs_bn_2d % 32) // 16)
)
b_scales = tl.load(b_scales_ptr + shuffle_idx)
if EVEN_K:
a_bf16 = tl.load(a_ptrs)
b = tl.load(b_ptrs, cache_modifier=cache_modifier)
else:
a_bf16 = tl.load(
a_ptrs,
mask=offs_k_bf16[None, :] < 2 * K,
other=0,
)
b = tl.load(
b_ptrs,
mask=offs_k[:, None] < K,
other=0,
cache_modifier=cache_modifier,
)
a, a_scales = _mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")
a_ptrs += BLOCK_SIZE_K * stride_ak
b_ptrs += (BLOCK_SIZE_K // 2) * stride_bk
ks_2d += BLOCK_SIZE_K // SCALE_GROUP_SIZE
c = accumulator.to(c_ptr.type.element_ty)
offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
c_ptrs = (
c_ptr
+ stride_cm * offs_cm[:, None]
+ stride_cn * offs_cn[None, :]
+ pid_k * stride_ck
)
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
if ATOMIC_ADD:
tl.atomic_add(c_ptrs, c, mask=c_mask, sem="relaxed")
else:
tl.store(c_ptrs, c, mask=c_mask)
_gemm_a16wfp4_shuf_repr = make_kernel_repr(
"_gemm_a16wfp4_shuffled_scales_kernel",
[
"BLOCK_SIZE_M",
"BLOCK_SIZE_N",
"BLOCK_SIZE_K",
"GROUP_SIZE_M",
"num_warps",
"num_stages",
"waves_per_eu",
"matrix_instr_nonkdim",
"cache_modifier",
"NUM_KSPLIT",
],
)
@triton.heuristics(
{
"EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)
and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),
"GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])
* triton.cdiv(args["N"], args["BLOCK_SIZE_N"]),
}
)
@triton.jit(repr=_gemm_a16wfp4_shuf_repr)
def _gemm_a16wfp4_shuffled_scales_kernel(
a_ptr,
b_ptr,
c_ptr,
b_scales_ptr,
M,
N,
K,
stride_am,
stride_ak,
stride_bk,
stride_bn,
stride_ck,
stride_cm,
stride_cn,
scale_sh_cols,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr,
NUM_KSPLIT: tl.constexpr,
SPLITK_BLOCK_SIZE: tl.constexpr,
EVEN_K: tl.constexpr,
num_warps: tl.constexpr,
num_stages: tl.constexpr,
waves_per_eu: tl.constexpr,
matrix_instr_nonkdim: tl.constexpr,
GRID_MN: tl.constexpr,
ATOMIC_ADD: tl.constexpr,
cache_modifier: tl.constexpr,
):
tl.assume(stride_am > 0)
tl.assume(stride_ak > 0)
tl.assume(stride_bk > 0)
tl.assume(stride_bn > 0)
tl.assume(stride_cm > 0)
tl.assume(stride_cn > 0)
tl.assume(scale_sh_cols > 0)
pid_unified = tl.program_id(axis=0)
pid_k = pid_unified % NUM_KSPLIT
pid = pid_unified // NUM_KSPLIT
num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
if NUM_KSPLIT == 1:
pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
else:
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
SCALE_GROUP_SIZE: tl.constexpr = 32
if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16
offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
a_ptrs = a_ptr + (
offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak
)
offs_k = tl.arange(0, BLOCK_SIZE_K // 2)
offs_k_split = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_k
offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
b_ptrs = b_ptr + (
offs_k_split[:, None] * stride_bk + offs_bn[None, :] * stride_bn
)
ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE)) + tl.arange(
0, BLOCK_SIZE_K // SCALE_GROUP_SIZE
)
offs_bn_2d = offs_bn[:, None]
ks_2d = ks[None, :]
sn8 = scale_sh_cols // 8
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
shuffle_idx = (
(offs_bn_2d // 32) * (sn8 * 256)
+ (ks_2d // 8) * 256
+ ((ks_2d % 8) % 4) * 64
+ (offs_bn_2d % 16) * 4
+ ((ks_2d % 8) // 4) * 2
+ ((offs_bn_2d % 32) // 16)
)
b_scales = tl.load(b_scales_ptr + shuffle_idx)
if EVEN_K:
a_bf16 = tl.load(a_ptrs)
b = tl.load(b_ptrs, cache_modifier=cache_modifier)
else:
a_bf16 = tl.load(
a_ptrs,
mask=offs_k_bf16[None, :] < 2 * K - k * BLOCK_SIZE_K,
other=0,
)
b = tl.load(
b_ptrs,
mask=offs_k[:, None] < K - k * (BLOCK_SIZE_K // 2),
other=0,
cache_modifier=cache_modifier,
)
a, a_scales = _mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")
a_ptrs += BLOCK_SIZE_K * stride_ak
b_ptrs += (BLOCK_SIZE_K // 2) * stride_bk
ks_2d += BLOCK_SIZE_K // SCALE_GROUP_SIZE
c = accumulator.to(c_ptr.type.element_ty)
offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
c_ptrs = (
c_ptr
+ stride_cm * offs_cm[:, None]
+ stride_cn * offs_cn[None, :]
+ pid_k * stride_ck
)
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
if ATOMIC_ADD:
tl.atomic_add(c_ptrs, c, mask=c_mask, sem="relaxed")
else:
tl.store(c_ptrs, c, mask=c_mask)
def _a16_shuf_get_config(M: int, N: int, K: int):
return get_gemm_config("GEMM-A16WFP4", M, N, 2 * K)
def _a16_shuf_get_splitk(K: int, BLOCK_SIZE_K: int, NUM_KSPLIT: int):
num_ksplit_step = 2
block_size_k_step = 2
splitk_block_size = (
triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K
)
while NUM_KSPLIT > 1 and BLOCK_SIZE_K > 16:
if (
K % (splitk_block_size // 2) == 0
and splitk_block_size % BLOCK_SIZE_K == 0
and K % (BLOCK_SIZE_K // 2) == 0
):
break
if K % (splitk_block_size // 2) != 0 and NUM_KSPLIT > 1:
NUM_KSPLIT = NUM_KSPLIT // num_ksplit_step
elif splitk_block_size % BLOCK_SIZE_K != 0:
if NUM_KSPLIT > 1:
NUM_KSPLIT = NUM_KSPLIT // num_ksplit_step
elif BLOCK_SIZE_K > 16:
BLOCK_SIZE_K = BLOCK_SIZE_K // block_size_k_step
elif K % (BLOCK_SIZE_K // 2) != 0 and BLOCK_SIZE_K > 16:
BLOCK_SIZE_K = BLOCK_SIZE_K // block_size_k_step
else:
break
splitk_block_size = (
triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K
)
return splitk_block_size, BLOCK_SIZE_K, NUM_KSPLIT
def gemm_a16wfp4_shuffled_scales(
x: torch.Tensor,
w: torch.Tensor,
w_scales_sh: torch.Tensor,
atomic_add: Optional[bool] = False,
dtype: Optional[torch.dtype] = torch.bfloat16,
y: Optional[torch.Tensor] = None,
config: Optional[str] = None,
) -> torch.Tensor:
assert arch_info.is_fp4_avail(), "MXFP4 is not available on your device"
M, K_bf16 = x.shape
N, K_packed = w.shape
assert K_packed * 2 == K_bf16, f"expected packed K/2 weights, got x={x.shape} w={w.shape}"
w = w.view(torch.uint8).T
w_scales_sh_u8 = w_scales_sh.view(torch.uint8).contiguous()
if config is None:
config, _ = _a16_shuf_get_config(M, N, K_packed)
else:
config = deserialize_str(config)
if y is None:
if atomic_add:
y = torch.zeros((M, N), dtype=dtype, device=x.device)
else:
y = torch.empty((M, N), dtype=dtype, device=x.device)
if config["NUM_KSPLIT"] > 1 and not atomic_add:
splitk_block_size, block_size_k, num_ksplit = _a16_shuf_get_splitk(
K_packed, config["BLOCK_SIZE_K"], config["NUM_KSPLIT"]
)
config["SPLITK_BLOCK_SIZE"] = splitk_block_size
config["BLOCK_SIZE_K"] = block_size_k
config["NUM_KSPLIT"] = num_ksplit
if config["BLOCK_SIZE_K"] >= 2 * K_packed:
config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * K_packed)
config["SPLITK_BLOCK_SIZE"] = 2 * K_packed
config["NUM_KSPLIT"] = 1
config["BLOCK_SIZE_K"] = max(config["BLOCK_SIZE_K"], 64)
if config["NUM_KSPLIT"] > 1 and not atomic_add:
y_pp = torch.empty(
(config["NUM_KSPLIT"], M, N), dtype=torch.float32, device=y.device
)
else:
config["SPLITK_BLOCK_SIZE"] = 2 * K_packed
y_pp = None
grid = lambda META: ( # noqa: E731
META["NUM_KSPLIT"]
* triton.cdiv(M, META["BLOCK_SIZE_M"])
* triton.cdiv(N, META["BLOCK_SIZE_N"]),
)
_gemm_a16wfp4_shuffled_scales_kernel[grid](
x,
w,
y if y_pp is None else y_pp,
w_scales_sh_u8.view(-1),
M,
N,
K_packed,
x.stride(0),
x.stride(1),
w.stride(0),
w.stride(1),
0 if y_pp is None else y_pp.stride(0),
y.stride(0) if y_pp is None else y_pp.stride(1),
y.stride(1) if y_pp is None else y_pp.stride(2),
w_scales_sh_u8.shape[1],
ATOMIC_ADD=atomic_add,
**config,
)
if config["NUM_KSPLIT"] > 1 and not atomic_add:
reduce_block_size_m = 16
reduce_block_size_n = 64
actual_ksplit = triton.cdiv(K_packed, (config["SPLITK_BLOCK_SIZE"] // 2))
grid_reduce = (
triton.cdiv(M, reduce_block_size_m),
triton.cdiv(N, reduce_block_size_n),
)
_gemm_afp4wfp4_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_size_m,
reduce_block_size_n,
actual_ksplit,
triton.next_power_of_2(config["NUM_KSPLIT"]),
)
return y
def _get_module():
global _module
from torch.utils.cpp_extension import load_inline
if _module is None:
_module = load_inline(
name="fused_quant_wave_combo_hybrid_v2",
cpp_sources=[CPP_SRC],
cuda_sources=[HIP_SRC],
functions=["fused_mxfp4_quant_out"],
verbose=False,
extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++17", "-O3"],
)
return _module
def _hip_quant_buffers(a: torch.Tensor):
m, k = a.shape
scale_n_valid = (k + 31) // 32
scale_n_pad = ((scale_n_valid + 7) // 8) * 8
scale_m_pad = ((m + 255) // 256) * 256
key = (a.device.index if a.device.index is not None else 0, m, k, scale_m_pad, scale_n_pad)
cached = _hip_buf_cache.get(key)
if cached is None:
a_fp4_u8 = torch.empty((m, k // 2), dtype=torch.uint8, device=a.device)
a_scale_u8 = torch.full((scale_m_pad * scale_n_pad,), 127, dtype=torch.uint8, device=a.device)
_hip_buf_cache[key] = (a_fp4_u8, a_scale_u8, scale_n_valid, scale_m_pad, scale_n_pad)
return _hip_buf_cache[key]
def _ensure_patch() -> None:
global _patched
if _patched:
return
import triton._utils as tu
tu.type_canonicalisation_dict["float4_e2m1fn_x2"] = "u8"
tu.type_canonicalisation_dict["float4_e2m1fn"] = "u8"
tu.type_canonicalisation_dict["float8_e8m0fnu"] = "u8"
tu.type_canonicalisation_dict["float8_e8m0fnuz"] = "u8"
tu.BITWIDTH_DICT["float4_e2m1fn_x2"] = 8
tu.BITWIDTH_DICT["float4_e2m1fn"] = 8
tu.BITWIDTH_DICT["float8_e8m0fnu"] = 8
tu.BITWIDTH_DICT["float8_e8m0fnuz"] = 8
_patched = True
def _unshuffle_e8m0(scale_sh: torch.Tensor, rows: int, cols: int) -> torch.Tensor:
sm, sn = scale_sh.shape
device_index = scale_sh.device.index if scale_sh.device.index is not None else 0
key = (device_index, sm, sn, rows, cols)
idx = _index_cache.get(key)
raw_flat = _raw_scale_buf_cache.get(key)
if idx is None or raw_flat is None:
# Precompute the inverse-shuffle gather plan once per shape. The tensor
# values change every ranked iteration, but the permutation does not.
sn8 = sn // 8
row_ids = torch.arange(rows, dtype=torch.int64).unsqueeze(1)
col_ids = torch.arange(cols, dtype=torch.int64).unsqueeze(0)
d0 = row_ids // 32
rem_r = row_ids % 32
d1 = rem_r // 16
d2 = rem_r % 16
d3 = col_ids // 8
rem_c = col_ids % 8
d4 = rem_c // 4
d5 = rem_c % 4
idx = (
d0 * (sn8 * 256)
+ d3 * 256
+ d5 * 64
+ d2 * 4
+ d4 * 2
+ d1
).reshape(-1).to(scale_sh.device)
raw_flat = torch.empty(rows * cols, dtype=torch.uint8, device=scale_sh.device)
_index_cache[key] = idx
_raw_scale_buf_cache[key] = raw_flat
scale_u8_flat = scale_sh.view(torch.uint8).reshape(-1)
torch.index_select(scale_u8_flat, 0, idx, out=raw_flat)
return raw_flat.view(rows, cols).view(scale_sh.dtype)
def _unshuffle_e8m0_view(scale_sh: torch.Tensor, rows: int, cols: int) -> torch.Tensor:
device_index = scale_sh.device.index if scale_sh.device.index is not None else 0
key = (id(scale_sh), scale_sh.data_ptr(), rows, cols, device_index)
cached = _view_scale_cache.get(key)
if cached is not None:
return cached[1]
sm, sn = scale_sh.shape
scale_u8 = scale_sh.view(torch.uint8).contiguous()
raw_u8 = (
scale_u8.view(sm // 32, sn // 8, 4, 16, 2, 2)
.permute(0, 5, 3, 1, 4, 2)
.contiguous()
.view(sm, sn)
)
raw = raw_u8[:rows, :cols].contiguous().view(scale_sh.dtype)
# Keep the source tensor alive so Python object ids cannot be recycled into
# stale cache hits across later ranked inputs.
_view_scale_cache[key] = (scale_sh, raw)
return raw
def _direct_output_buffer(a: torch.Tensor, n: int) -> torch.Tensor:
key = (a.device.index, a.shape[0], n)
cached = _direct_out_cache.get(key)
if cached is None:
cached = torch.empty((a.shape[0], n), dtype=torch.bfloat16, device=a.device)
_direct_out_cache[key] = cached
return cached
def _fused_pp_buffer(device: torch.device, num_ksplit: int, m: int, n: int) -> torch.Tensor:
device_index = device.index if device.index is not None else 0
key = (device_index, num_ksplit, m, n)
cached = _fused_pp_cache.get(key)
if cached is None:
cached = torch.empty((num_ksplit, m, n), dtype=torch.float32, device=device)
_fused_pp_cache[key] = cached
return cached
def gemm_a16wfp4_fused(
A: torch.Tensor,
B_q: torch.Tensor,
B_scale_sh: torch.Tensor,
config_override: dict | None = None,
atomic_add: bool = False,
y: torch.Tensor | None = None,
):
M, K_bf16 = A.shape
B_q_u8 = B_q.view(torch.uint8).contiguous()
B_scale_sh_u8 = B_scale_sh.view(torch.uint8).contiguous()
N, K_packed = B_q_u8.shape
if config_override is None:
config, _ = get_gemm_config("GEMM-A16WFP4", M, N, K_bf16)
else:
config = dict(config_override)
if config["NUM_KSPLIT"] > 1 and not atomic_add:
splitk_block_size, block_size_k, num_ksplit = _a16_shuf_get_splitk(
K_packed, config["BLOCK_SIZE_K"], config["NUM_KSPLIT"]
)
config["SPLITK_BLOCK_SIZE"] = splitk_block_size
config["BLOCK_SIZE_K"] = block_size_k
config["NUM_KSPLIT"] = num_ksplit
if config["BLOCK_SIZE_K"] >= 2 * K_packed:
config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * K_packed)
config["SPLITK_BLOCK_SIZE"] = 2 * K_packed
config["NUM_KSPLIT"] = 1
config["BLOCK_SIZE_K"] = max(config["BLOCK_SIZE_K"], 64)
if y is None:
y = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
if config["NUM_KSPLIT"] > 1 and not atomic_add:
y_pp = _fused_pp_buffer(A.device, config["NUM_KSPLIT"], M, N)
else:
y_pp = None
grid = lambda META: ( # noqa: E731
META["NUM_KSPLIT"]
* triton.cdiv(M, META["BLOCK_SIZE_M"])
* triton.cdiv(N, META["BLOCK_SIZE_N"]),
)
B = B_q_u8.t()
_gemm_a16wfp4_fused_kernel[grid](
A,
B,
y if y_pp is None else y_pp,
B_scale_sh_u8.view(-1),
M,
N,
K_packed,
A.stride(0),
A.stride(1),
B.stride(0),
B.stride(1),
0 if y_pp is None else y_pp.stride(0),
y.stride(0) if y_pp is None else y_pp.stride(1),
y.stride(1) if y_pp is None else y_pp.stride(2),
B_scale_sh_u8.shape[1],
ATOMIC_ADD=atomic_add,
**config,
)
if config["NUM_KSPLIT"] > 1 and not atomic_add:
reduce_block_size_m = 16
reduce_block_size_n = 64
actual_ksplit = triton.cdiv(K_packed, (config["SPLITK_BLOCK_SIZE"] // 2))
grid_reduce = (
triton.cdiv(M, reduce_block_size_m),
triton.cdiv(N, reduce_block_size_n),
)
_gemm_afp4wfp4_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_size_m,
reduce_block_size_n,
actual_ksplit,
triton.next_power_of_2(config["NUM_KSPLIT"]),
)
return y
def _direct_path(A: torch.Tensor, B: torch.Tensor, B_q: torch.Tensor, B_scale_sh: torch.Tensor):
n = B.shape[0]
out = _direct_output_buffer(A, n)
return gemm_a16wfp4_shuffled_scales(A, B_q, B_scale_sh, dtype=torch.bfloat16, y=out)
def _direct_raw_outlier_path(A: torch.Tensor, B: torch.Tensor, B_q: torch.Tensor, B_scale_sh: torch.Tensor):
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4
n = B.shape[0]
k = A.shape[1]
_ensure_patch()
out = _direct_output_buffer(A, n)
B_scale_raw = fast_unshuffle_e8m0(B_scale_sh, n, k // 32)
return gemm_a16wfp4(
A,
B_q.view(torch.uint8),
B_scale_raw,
dtype=torch.bfloat16,
y=out,
config=OUTLIER_RAW_CONFIG,
)
def _direct_fused_outlier_path(A: torch.Tensor, B: torch.Tensor, B_q: torch.Tensor, B_scale_sh: torch.Tensor):
n = B.shape[0]
out = _direct_output_buffer(A, n)
return gemm_a16wfp4_fused(
A,
B_q,
B_scale_sh,
config_override=OUTLIER_RAW_CONFIG,
atomic_add=False,
y=out,
)
def _hip_path(a: torch.Tensor, b: torch.Tensor, b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor):
m, k = a.shape
n = b.shape[0]
try:
mod = _get_module()
a_fp4_u8, a_scale_u8, scale_n_valid, scale_m_pad, scale_n_pad = _hip_quant_buffers(a)
mod.fused_mxfp4_quant_out(a, a_fp4_u8, a_scale_u8, scale_n_valid, scale_n_pad)
a_q = a_fp4_u8.view(dtypes.fp4x2)
a_scale_sh = a_scale_u8.view(scale_m_pad, scale_n_pad).view(dtypes.fp8_e8m0)
except Exception as exc:
_status((m, n, k, "hip-fallback"), f"HIP quant failed for {(m, n, k)}: {repr(exc)}")
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
x_fp4, bs_e8m0 = dynamic_mxfp4_quant(a)
a_q = x_fp4.view(dtypes.fp4x2)
a_scale_sh = e8m0_shuffle(bs_e8m0)
config = TUNED_CONFIGS.get((m, n, k))
if config is not None:
knl_name, split_k = config
padded_m = get_padded_m(m, n, k, 0)
out = torch.empty(padded_m, n, dtype=torch.bfloat16, device=a.device)
gemm_a4w4_asm(
a_q, b_shuffle, a_scale_sh, b_scale_sh,
out, knl_name, None, 1.0, 0.0, True, split_k,
)
return out[:m]
return aiter.gemm_a4w4(
a_q, b_shuffle, a_scale_sh, b_scale_sh,
dtype=dtypes.bf16, bpreshuffle=True,
)
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
shape = (A.shape[0], B.shape[0], A.shape[1])
_ensure_patch()
if shape == OUTLIER_RAW_SHAPE:
_status((shape, "route"), f"route fused-outlier for {shape}")
try:
return _direct_fused_outlier_path(A, B, B_q, B_scale_sh)
except Exception as exc:
_status((shape, "direct-fused-fail"), f"fused outlier failed for {shape}: {repr(exc)}")
try:
return _direct_raw_outlier_path(A, B, B_q, B_scale_sh)
except Exception as exc:
_status((shape, "direct-raw-fail"), f"direct raw custom failed for {shape}: {repr(exc)}")
if shape in DIRECT_SHAPES:
_status((shape, "route"), f"route direct for {shape}")
try:
return _direct_path(A, B, B_q, B_scale_sh)
except Exception as exc:
_status((shape, "direct-fail"), f"direct failed for {shape}: {repr(exc)}")
_status((shape, "route"), f"route HIP for {shape}")
return _hip_path(A, B, B_shuffle, B_scale_sh)
scrolls · 1082 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