submission 615003
trulyspinach · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 905 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-615003?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:2e3159a6fdcffc06e707786c1d52af48db99ec480c3ff86d65a432fbd31e68a6
license declaredunknown
license concludedunknown
authorstrulyspinach
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
shared-memory
__shared__ float absmax[ROWS_PER_BLOCK * PAIRS_PER_ROW];split-k
def _get_splitk(k_packed: int, block_size_k: int, num_ksplit: int):Kernel source
submission.py905 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
import os
os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
os.environ.setdefault("CXX", "clang++")
from task import input_t, output_t
import torch
import triton
import triton.language as tl
from torch.utils.cpp_extension import load_inline
from aiter import dtypes
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid, remap_xcd
from aiter.ops.triton.utils.gemm_config_utils import get_gemm_config
from aiter.utility.fp4_utils import e8m0_shuffle
_M_BOUNDS = (4, 8, 16, 31, 32, 64, 128, 256, 512, 1024, 2048, 4096, 8192)
_REFERENCE_SHAPES: set[tuple[int, int, int]] = {(64, 7168, 2048), (256, 3072, 1536)}
_B_TRITON_CACHE: dict[tuple[int, int, int, tuple[int, ...], tuple[int, ...]], tuple[torch.Tensor, torch.Tensor]] = {}
_FUSED_CONFIGS = {
(2880, 512): {
"M_LEQ_8": {
"BLOCK_SIZE_M": 8,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 4,
"num_warps": 2,
"num_stages": 2,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
},
"M_LEQ_31": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 4,
"num_warps": 2,
"num_stages": 2,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
},
"M_LEQ_32": {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 4,
"num_warps": 2,
"num_stages": 2,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
},
},
(4096, 512): {
"M_LEQ_32": {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 1,
},
},
}
_OVERRIDE_CONFIGS = {
(2112, 7168): {
"M_LEQ_31": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 1024,
"GROUP_SIZE_M": 1,
"num_warps": 2,
"num_stages": 2,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
},
}
CPP_SRC = r"""
void quant_mxfp4(torch::Tensor input, torch::Tensor out_q, torch::Tensor out_scale);
void quant_mxfp4_shuffled(torch::Tensor input, torch::Tensor out_q, torch::Tensor out_scale);
void shuffle_e8m0_scales(torch::Tensor input, torch::Tensor output);
"""
CUDA_SRC = r"""
#include <cmath>
#include <cstdint>
#include <stdexcept>
#include <hip/hip_runtime.h>
#include <torch/extension.h>
__device__ __forceinline__ float bf16_to_float(uint16_t x) {
uint32_t bits = static_cast<uint32_t>(x) << 16;
return __uint_as_float(bits);
}
__device__ __forceinline__ uint8_t float_to_mxfp4(float x) {
uint32_t qx = __float_as_uint(x);
uint32_t s = qx & 0x80000000u;
qx ^= s;
float qx_fp32 = __uint_as_float(qx);
bool saturate_mask = qx_fp32 >= 6.0f;
bool denormal_mask = (!saturate_mask) && (qx_fp32 < 1.0f);
constexpr uint32_t MBITS_F32 = 23u;
constexpr uint32_t MBITS_FP4 = 1u;
constexpr uint32_t EXP_BIAS_FP32 = 127u;
constexpr uint32_t EXP_BIAS_FP4 = 1u;
constexpr uint32_t denorm_mask_int =
((EXP_BIAS_FP32 - EXP_BIAS_FP4) + (MBITS_F32 - MBITS_FP4) + 1u) << MBITS_F32;
float denorm_mask_float = __uint_as_float(denorm_mask_int);
uint8_t denormal_x = 0;
if (denormal_mask) {
uint32_t denormal_bits = __float_as_uint(qx_fp32 + denorm_mask_float);
denormal_x = static_cast<uint8_t>(denormal_bits - denorm_mask_int);
}
uint8_t normal_x = 0;
if (!saturate_mask && !denormal_mask) {
uint32_t mant_odd = (qx >> (MBITS_F32 - MBITS_FP4)) & 1u;
uint32_t val_to_add =
((EXP_BIAS_FP4 - EXP_BIAS_FP32) << MBITS_F32) + (1u << 21) - 1u;
uint32_t rounded = qx + val_to_add + mant_odd;
normal_x = static_cast<uint8_t>(rounded >> (MBITS_F32 - MBITS_FP4));
}
uint8_t value = saturate_mask ? 0x7u : normal_x;
value = denormal_mask ? denormal_x : value;
value |= static_cast<uint8_t>(s >> 28);
return value;
}
template <bool SHUFFLED_SCALE>
__global__ void quant_mxfp4_kernel(
const uint16_t* input,
uint8_t* out_q,
uint8_t* out_scale,
int M,
int N,
int stride_in_m,
int stride_in_n,
int stride_q_m,
int stride_q_n,
int stride_s_m,
int stride_s_n) {
constexpr int ROWS_PER_BLOCK = 16;
constexpr int PAIRS_PER_ROW = 16;
constexpr int GROUP_SIZE = 32;
int tid = threadIdx.x;
int row_local = tid / PAIRS_PER_ROW;
int pair_idx = tid % PAIRS_PER_ROW;
int row = blockIdx.x * ROWS_PER_BLOCK + row_local;
int group = blockIdx.y;
int groups = N / GROUP_SIZE;
__shared__ float absmax[ROWS_PER_BLOCK * PAIRS_PER_ROW];
__shared__ float quant_scale[ROWS_PER_BLOCK];
bool active = row < M && group < groups;
int col0 = group * GROUP_SIZE + pair_idx * 2;
int col1 = col0 + 1;
float x0 = 0.0f;
float x1 = 0.0f;
if (active) {
x0 = bf16_to_float(input[row * stride_in_m + col0 * stride_in_n]);
x1 = bf16_to_float(input[row * stride_in_m + col1 * stride_in_n]);
}
int shared_idx = row_local * PAIRS_PER_ROW + pair_idx;
absmax[shared_idx] = fmaxf(fabsf(x0), fabsf(x1));
__syncthreads();
for (int offset = PAIRS_PER_ROW / 2; offset > 0; offset >>= 1) {
if (pair_idx < offset) {
absmax[shared_idx] = fmaxf(absmax[shared_idx], absmax[shared_idx + offset]);
}
__syncthreads();
}
if (pair_idx == 0) {
float qscale = 0.0f;
uint8_t scale_byte = 0;
if (active) {
float amax = absmax[row_local * PAIRS_PER_ROW];
uint32_t rounded_bits = (__float_as_uint(amax) + 0x200000u) & 0xFF800000u;
if (rounded_bits == 0u) {
scale_byte = 0;
qscale = __uint_as_float(254u << 23);
} else {
int exponent = static_cast<int>((rounded_bits >> 23) & 0xFFu) - 127;
int scale_unbiased = exponent - 2;
scale_unbiased = scale_unbiased < -127 ? -127 : scale_unbiased;
scale_unbiased = scale_unbiased > 127 ? 127 : scale_unbiased;
scale_byte = static_cast<uint8_t>(scale_unbiased + 127);
qscale = __uint_as_float(static_cast<uint32_t>(127 - scale_unbiased) << 23);
}
if constexpr (SHUFFLED_SCALE) {
int row_block = row / 32;
int row_half = (row % 32) / 16;
int row_inner = row % 16;
int group_block = group / 8;
int group_half = (group % 8) / 4;
int group_inner = group % 4;
int out_cols = stride_s_m;
int rest =
group_block * 256 + group_inner * 64 + row_inner * 4 + group_half * 2 + row_half;
int out_row = row_block * 32 + rest / out_cols;
int out_col = rest % out_cols;
out_scale[out_row * stride_s_m + out_col * stride_s_n] = scale_byte;
} else {
out_scale[row * stride_s_m + group * stride_s_n] = scale_byte;
}
}
quant_scale[row_local] = qscale;
}
__syncthreads();
if (!active) {
return;
}
uint8_t q0 = float_to_mxfp4(x0 * quant_scale[row_local]);
uint8_t q1 = float_to_mxfp4(x1 * quant_scale[row_local]);
out_q[row * stride_q_m + (group * PAIRS_PER_ROW + pair_idx) * stride_q_n] =
static_cast<uint8_t>(q0 | (q1 << 4));
}
void quant_mxfp4(torch::Tensor input, torch::Tensor out_q, torch::Tensor out_scale) {
int M = input.size(0);
int N = input.size(1);
dim3 blocks((M + 15) / 16, N / 32);
dim3 threads(256);
quant_mxfp4_kernel<false><<<blocks, threads>>>(
reinterpret_cast<const uint16_t*>(input.data_ptr()),
out_q.data_ptr<uint8_t>(),
out_scale.data_ptr<uint8_t>(),
M,
N,
input.stride(0),
input.stride(1),
out_q.stride(0),
out_q.stride(1),
out_scale.stride(0),
out_scale.stride(1)
);
hipError_t err = hipGetLastError();
if (err != hipSuccess) {
throw std::runtime_error(hipGetErrorString(err));
}
}
void quant_mxfp4_shuffled(torch::Tensor input, torch::Tensor out_q, torch::Tensor out_scale) {
int M = input.size(0);
int N = input.size(1);
dim3 blocks((M + 15) / 16, N / 32);
dim3 threads(256);
quant_mxfp4_kernel<true><<<blocks, threads>>>(
reinterpret_cast<const uint16_t*>(input.data_ptr()),
out_q.data_ptr<uint8_t>(),
out_scale.data_ptr<uint8_t>(),
M,
N,
input.stride(0),
input.stride(1),
out_q.stride(0),
out_q.stride(1),
out_scale.stride(0),
out_scale.stride(1)
);
hipError_t err = hipGetLastError();
if (err != hipSuccess) {
throw std::runtime_error(hipGetErrorString(err));
}
}
__global__ void shuffle_e8m0_scales_kernel(
const uint8_t* input,
uint8_t* output,
int M,
int G,
int stride_in_m,
int stride_in_g,
int stride_out_m,
int stride_out_g) {
int row = blockIdx.x;
int group = blockIdx.y * blockDim.x + threadIdx.x;
if (row >= M || group >= G) {
return;
}
int row_block = row / 32;
int row_half = (row % 32) / 16;
int row_inner = row % 16;
int group_block = group / 8;
int group_half = (group % 8) / 4;
int group_inner = group % 4;
int out_cols = stride_out_m;
int rest = group_block * 256 + group_inner * 64 + row_inner * 4 + group_half * 2 + row_half;
int out_row = row_block * 32 + rest / out_cols;
int out_col = rest % out_cols;
output[out_row * stride_out_m + out_col * stride_out_g] =
input[row * stride_in_m + group * stride_in_g];
}
void shuffle_e8m0_scales(torch::Tensor input, torch::Tensor output) {
int M = input.size(0);
int G = input.size(1);
dim3 blocks(M, (G + 127) / 128);
dim3 threads(128);
shuffle_e8m0_scales_kernel<<<blocks, threads>>>(
input.data_ptr<uint8_t>(),
output.data_ptr<uint8_t>(),
M,
G,
input.stride(0),
input.stride(1),
output.stride(0),
output.stride(1)
);
hipError_t err = hipGetLastError();
if (err != hipSuccess) {
throw std::runtime_error(hipGetErrorString(err));
}
}
"""
_INLINE_QUANT = load_inline(
name="submission_inline_quant_v1",
cpp_sources=[CPP_SRC],
cuda_sources=[CUDA_SRC],
functions=["quant_mxfp4", "quant_mxfp4_shuffled", "shuffle_e8m0_scales"],
verbose=False,
extra_cuda_cflags=["--offload-arch=gfx950", "-std=c++20"],
)
def _tensor_token(x: torch.Tensor) -> tuple[int, int]:
return (id(x), int(x._version))
def _tensor_cache_key(x: torch.Tensor) -> tuple[int, int, int]:
token_id, token_version = _tensor_token(x)
return (token_id, token_version, int(x.data_ptr()))
def _inline_dynamic_mxfp4_quant(a: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
m, k = int(a.shape[0]), int(a.shape[1])
a_q = torch.empty((m, k // 2), dtype=torch.uint8, device=a.device)
a_scale = torch.empty((m, k // 32), dtype=torch.uint8, device=a.device)
_INLINE_QUANT.quant_mxfp4(a, a_q, a_scale)
return a_q, a_scale.view(dtypes.fp8_e8m0)
def _inline_dynamic_mxfp4_quant_shuffled(a: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
m, k = int(a.shape[0]), int(a.shape[1])
a_q = torch.empty((m, k // 2), dtype=torch.uint8, device=a.device)
out_m = ((m + 255) // 256) * 256
out_g = (((k // 32) + 7) // 8) * 8
a_scale_sh = torch.zeros((out_m, out_g), dtype=torch.uint8, device=a.device)
_INLINE_QUANT.quant_mxfp4_shuffled(a, a_q, a_scale_sh)
return a_q, a_scale_sh.view(dtypes.fp8_e8m0)
def _inline_e8m0_shuffle(a_scale: torch.Tensor) -> torch.Tensor:
m, g = int(a_scale.shape[0]), int(a_scale.shape[1])
out_m = ((m + 255) // 256) * 256
out_g = ((g + 7) // 8) * 8
shuffled = torch.zeros((out_m, out_g), dtype=torch.uint8, device=a_scale.device)
_INLINE_QUANT.shuffle_e8m0_scales(a_scale.view(torch.uint8), shuffled)
return shuffled.view(dtypes.fp8_e8m0)
def _select_config(configs: dict[tuple[int, int], dict], m: int, n: int, k: int) -> dict | None:
override = configs.get((n, k))
if override is None:
return None
for bound in _M_BOUNDS:
key = f"M_LEQ_{bound}"
if m <= bound and key in override:
return dict(override[key])
if "any" in override:
return dict(override["any"])
return None
def _get_config(m: int, n: int, k_packed: int) -> dict:
override = _select_config(_OVERRIDE_CONFIGS, m, n, 2 * k_packed)
if override is not None:
return override
config, _ = get_gemm_config(
"GEMM-AFP4WFP4_PRESHUFFLED",
m,
n,
2 * k_packed,
bounds=_M_BOUNDS,
)
return config
def _get_splitk(k_packed: int, block_size_k: int, num_ksplit: int):
splitk_block_size = (
triton.cdiv((2 * triton.cdiv(k_packed, num_ksplit)), block_size_k) * block_size_k
)
while num_ksplit > 1 and block_size_k > 16:
if (
k_packed % (splitk_block_size // 2) == 0
and splitk_block_size % block_size_k == 0
and k_packed % (block_size_k // 2) == 0
):
break
if k_packed % (splitk_block_size // 2) != 0 and num_ksplit > 1:
num_ksplit //= 2
elif splitk_block_size % block_size_k != 0:
if num_ksplit > 1:
num_ksplit //= 2
elif block_size_k > 16:
block_size_k //= 2
elif k_packed % (block_size_k // 2) != 0 and block_size_k > 16:
block_size_k //= 2
else:
break
splitk_block_size = (
triton.cdiv((2 * triton.cdiv(k_packed, num_ksplit)), block_size_k)
* block_size_k
)
num_ksplit = triton.cdiv(k_packed, (splitk_block_size // 2))
return splitk_block_size, block_size_k, num_ksplit
@triton.heuristics(
{
"EVEN_K": lambda args: (args["k_packed"] % (args["BLOCK_SIZE_K"] // 2) == 0)
and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
and (args["k_packed"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),
}
)
@triton.jit
def _gemm_afp4wfp4_preshuffle_kernel(
a_ptr,
b_ptr,
c_ptr,
a_scales_ptr,
b_scales_ptr,
m,
n,
k_packed,
stride_am,
stride_ak,
stride_bn,
stride_bk,
stride_ck,
stride_cm,
stride_cn,
stride_asm,
stride_ask,
stride_bsn,
stride_bsk,
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,
cache_modifier: tl.constexpr,
):
tl.assume(stride_am > 0)
tl.assume(stride_ak > 0)
tl.assume(stride_bn > 0)
tl.assume(stride_bk > 0)
tl.assume(stride_cm > 0)
tl.assume(stride_cn > 0)
tl.assume(stride_asm > 0)
tl.assume(stride_ask > 0)
tl.assume(stride_bsn > 0)
tl.assume(stride_bsk > 0)
scale_group_size: tl.constexpr = 32
grid_mn = tl.cdiv(m, BLOCK_SIZE_M) * tl.cdiv(n, BLOCK_SIZE_N)
pid_unified = tl.program_id(axis=0)
pid_unified = remap_xcd(pid_unified, grid_mn * NUM_KSPLIT, NUM_XCDS=8)
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
if (pid_k * SPLITK_BLOCK_SIZE // 2) < k_packed:
num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
offs_k = tl.arange(0, BLOCK_SIZE_K // 2)
offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
offs_k_split = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_k
offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % m
offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % n
a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_k_split[None, :] * stride_ak
b_ptrs = b_ptr + offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk
offs_asn = (
pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)
) % n
offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // scale_group_size) * 32) + tl.arange(
0, BLOCK_SIZE_K // scale_group_size * 32
)
b_scale_ptrs = (
b_scales_ptr
+ offs_asn[:, None] * stride_bsn
+ offs_ks[None, :] * stride_bsk
)
if BLOCK_SIZE_M < 32:
offs_ks_non_shufl = (
pid_k * (SPLITK_BLOCK_SIZE // scale_group_size)
) + tl.arange(0, BLOCK_SIZE_K // scale_group_size)
a_scale_ptrs = (
a_scales_ptr
+ offs_am[:, None] * stride_asm
+ offs_ks_non_shufl[None, :] * stride_ask
)
else:
offs_asm = (
pid_m * (BLOCK_SIZE_M // 32) + tl.arange(0, BLOCK_SIZE_M // 32)
) % m
a_scale_ptrs = (
a_scales_ptr
+ offs_asm[:, None] * stride_asm
+ offs_ks[None, :] * stride_ask
)
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k_iter in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
if BLOCK_SIZE_M < 32:
a_scales = tl.load(a_scale_ptrs)
else:
a_scales = (
tl.load(a_scale_ptrs)
.reshape(
BLOCK_SIZE_M // 32,
BLOCK_SIZE_K // scale_group_size // 8,
4,
16,
2,
2,
1,
)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // scale_group_size)
)
b_scales = (
tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
.reshape(
BLOCK_SIZE_N // 32,
BLOCK_SIZE_K // scale_group_size // 8,
4,
16,
2,
2,
1,
)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // scale_group_size)
)
if EVEN_K:
a = tl.load(a_ptrs)
b = tl.load(b_ptrs, cache_modifier=cache_modifier)
else:
k_remaining = k_packed - k_iter * (BLOCK_SIZE_K // 2)
a = tl.load(a_ptrs, mask=offs_k[None, :] < k_remaining, other=0)
b = tl.load(
b_ptrs,
mask=offs_k_shuffle_arr[None, :] < (k_remaining * 16),
other=0,
cache_modifier=cache_modifier,
)
b = (
b.reshape(
1,
BLOCK_SIZE_N // 16,
BLOCK_SIZE_K // 64,
2,
16,
16,
)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
.trans(1, 0)
)
accumulator = tl.dot_scaled(
a,
a_scales,
"e2m1",
b,
b_scales,
"e2m1",
accumulator,
)
a_ptrs += (BLOCK_SIZE_K // 2) * stride_ak
b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
if BLOCK_SIZE_M < 32:
a_scale_ptrs += (BLOCK_SIZE_K // scale_group_size) * stride_ask
else:
a_scale_ptrs += BLOCK_SIZE_K * stride_ask
b_scale_ptrs += BLOCK_SIZE_K * stride_bsk
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)
tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")
@triton.jit
def _reduce_kernel(
c_in_ptr,
c_out_ptr,
m,
n,
stride_c_in_k,
stride_c_in_m,
stride_c_in_n,
stride_c_out_m,
stride_c_out_n,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
ACTUAL_KSPLIT: tl.constexpr,
MAX_KSPLIT: tl.constexpr,
):
pid_m = tl.program_id(axis=0)
pid_n = tl.program_id(axis=1)
offs_m = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % m
offs_n = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % n
offs_k = tl.arange(0, MAX_KSPLIT)
c_in_ptrs = (
c_in_ptr
+ offs_k[:, None, None] * stride_c_in_k
+ offs_m[None, :, None] * stride_c_in_m
+ offs_n[None, None, :] * stride_c_in_n
)
if ACTUAL_KSPLIT == MAX_KSPLIT:
c = tl.load(c_in_ptrs)
else:
c = tl.load(c_in_ptrs, mask=offs_k[:, None, None] < ACTUAL_KSPLIT)
c = tl.sum(c, axis=0)
c = c.to(c_out_ptr.type.element_ty)
c_out_ptrs = c_out_ptr + offs_m[:, None] * stride_c_out_m + offs_n[None, :] * stride_c_out_n
tl.store(c_out_ptrs, c)
def _run_custom_preshuffle(
a_q_u8: torch.Tensor,
b_shuffled_u8: torch.Tensor,
a_scale_triton_u8: torch.Tensor,
b_scale_triton_u8: torch.Tensor,
dtype: torch.dtype = torch.bfloat16,
config_override: dict | None = None,
) -> torch.Tensor:
m, k_packed = a_q_u8.shape
n_16, k_16 = b_shuffled_u8.shape
n = n_16 * 16
assert k_16 == k_packed * 16, (a_q_u8.shape, b_shuffled_u8.shape)
config = dict(config_override) if config_override is not None else _get_config(m, n, k_packed)
if config["NUM_KSPLIT"] > 1:
splitk_block_size, block_size_k, num_ksplit = _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
else:
config["SPLITK_BLOCK_SIZE"] = 2 * k_packed
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["BLOCK_SIZE_N"] = max(config["BLOCK_SIZE_N"], 32)
if m < 32:
assert config["BLOCK_SIZE_M"] <= 16
else:
assert config["BLOCK_SIZE_M"] >= 32
y = torch.empty((m, n), dtype=dtype, device=a_q_u8.device)
y_pp = None
if config["NUM_KSPLIT"] > 1:
y_pp = torch.empty((config["NUM_KSPLIT"], m, n), dtype=torch.float32, device=a_q_u8.device)
grid = lambda meta: ( # noqa: E731
(
meta["NUM_KSPLIT"]
* triton.cdiv(m, meta["BLOCK_SIZE_M"])
* triton.cdiv(n, meta["BLOCK_SIZE_N"])
),
)
_gemm_afp4wfp4_preshuffle_kernel[grid](
a_q_u8,
b_shuffled_u8,
y if y_pp is None else y_pp,
a_scale_triton_u8,
b_scale_triton_u8,
m,
n,
k_packed,
a_q_u8.stride(0),
a_q_u8.stride(1),
b_shuffled_u8.stride(0),
b_shuffled_u8.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),
a_scale_triton_u8.stride(0),
a_scale_triton_u8.stride(1),
b_scale_triton_u8.stride(0),
b_scale_triton_u8.stride(1),
**config,
)
if y_pp is None:
return y
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),
)
_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 _prepare_a_scale(a_scale: torch.Tensor, m: int) -> torch.Tensor:
if m < 32:
return a_scale.view(torch.uint8).contiguous()
return _inline_e8m0_shuffle(a_scale).view(torch.uint8).reshape(-1, a_scale.shape[1] * 32).contiguous()
def _prepare_b_for_triton(b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor):
key = (
int(b_shuffle.device.index or 0),
*_tensor_cache_key(b_shuffle),
*_tensor_cache_key(b_scale_sh),
tuple(int(dim) for dim in b_shuffle.shape),
tuple(int(dim) for dim in b_scale_sh.shape),
)
cached = _B_TRITON_CACHE.get(key)
if cached is not None:
return cached
b_shuffle_u8 = b_shuffle.view(torch.uint8).reshape(b_shuffle.shape[0] // 16, b_shuffle.shape[1] * 16)
b_scale_u8 = b_scale_sh.view(torch.uint8).reshape(b_scale_sh.shape[0] // 32, b_scale_sh.shape[1] * 32)
cached = (b_shuffle_u8.contiguous(), b_scale_u8.contiguous())
_B_TRITON_CACHE[key] = cached
return cached
def _run_fused_preshuffle(
a: torch.Tensor,
b_shuffled_u8: torch.Tensor,
b_scale_triton_u8: torch.Tensor,
config: dict,
):
return gemm_a16wfp4_preshuffle(
a,
b_shuffled_u8,
b_scale_triton_u8,
dtype=dtypes.bf16,
config=config,
)
def _compute_uncached(
a: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
) -> torch.Tensor:
shape = (int(a.shape[0]), int(b_shuffle.shape[0]), int(a.shape[1]))
try:
b_shuffle_triton, b_scale_triton = _prepare_b_for_triton(b_shuffle, b_scale_sh)
fused_config = _select_config(_FUSED_CONFIGS, *shape)
if fused_config is not None:
return _run_fused_preshuffle(
a,
b_shuffle_triton,
b_scale_triton,
fused_config,
)
if shape in _REFERENCE_SHAPES:
a_q_u8, a_scale_sh = _inline_dynamic_mxfp4_quant_shuffled(a)
return _reference(a_q_u8.view(dtypes.fp4x2), b_shuffle, a_scale_sh, b_scale_sh)
a_q_u8, a_scale = _inline_dynamic_mxfp4_quant(a)
if shape == (16, 2112, 7168):
return _run_custom_preshuffle(
a_q_u8,
b_shuffle_triton,
_prepare_a_scale(a_scale, a.shape[0]),
b_scale_triton,
dtype=torch.bfloat16,
config_override=_OVERRIDE_CONFIGS[(2112, 7168)]["M_LEQ_31"],
)
return _run_custom_preshuffle(
a_q_u8,
b_shuffle_triton,
_prepare_a_scale(a_scale, a.shape[0]),
b_scale_triton,
dtype=torch.bfloat16,
)
except Exception:
a_q, a_scale = dynamic_mxfp4_quant(a)
a_scale_sh = e8m0_shuffle(a_scale)
return _reference(a_q, b_shuffle, a_scale_sh, b_scale_sh)
def _reference(a_q: torch.Tensor, b_shuffle: torch.Tensor, a_scale_sh: torch.Tensor, b_scale_sh: torch.Tensor):
import aiter
return aiter.gemm_a4w4(
a_q.view(dtypes.fp4x2),
b_shuffle,
a_scale_sh.view(dtypes.fp8_e8m0),
b_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
def custom_kernel(data: input_t) -> output_t:
a, _b, _b_q, b_shuffle, b_scale_sh = data
a = a.contiguous()
b_shuffle = b_shuffle.contiguous()
b_scale_sh = b_scale_sh.contiguous()
return _compute_uncached(a, b_shuffle, b_scale_sh)
scrolls · 905 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