submission 755023
Shaw · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 721 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-755023?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:346e55eaa3c65d6f5781e0d93924dd7bf9108789f9844703b091a0bdd4f22129
license declaredunknown
license concludedunknown
authorsShaw
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
Converts x (fp32) to mxfp4 format.num-warps = 1
NUM_WARPS = 1persistent-kernel
b_pid = (pid_m - N_QUANT_M) * tl.num_programs(1) + tl.program_id(1)split-k
def _splitk_fused_quant_gemm_kernel(stages = 1
NUM_STAGES = 1tile-k = 512
def _run_splitk_fused_quant_gemm(A, B_shuffle, B_scale_sh, SPLIT_K=8, BLOCK_M=16, BLOCK_K=512):tile-m = 64
BLOCK_SIZE_M = 64tile-n = 32
BLOCK_SIZE_N = 32Kernel source
submission.py721 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
EXP-20260330-C: Triton codegen optimizations.
Integer bitops in quant (replace tl.log2/tl.exp2), fast_math on dot_scaled,
cache_modifier=".cg" on B loads.
Baseline: EXP-20260328-17 @ 10.86 us, rank 66.
"""
import torch
import triton
import triton.language as tl
import aiter
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
from task import input_t, output_t
import os
SCALE_GROUP_SIZE = 32
# ========== Forked _mxfp4_quant_op with integer bitops ==========
@triton.jit
def _mxfp4_quant_op(
x,
BLOCK_SIZE_N,
BLOCK_SIZE_M,
MXFP4_QUANT_BLOCK_SIZE,
):
"""
Converts x (fp32) to mxfp4 format.
Forked from aiter with tl.log2/tl.exp2 replaced by integer bit extraction.
"""
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)
# Calculate scale
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
# Integer bitops: extract exponent directly instead of tl.log2 + tl.exp2
exponent = ((amax >> 23) & 0xFF).to(tl.int32)
scale_e8m0_unbiased = exponent - 129 # 127 (IEEE bias) + 2
# tl.clamp doesn't support int32, use tl.where instead
scale_e8m0_unbiased = tl.where(scale_e8m0_unbiased < -127, -127, scale_e8m0_unbiased)
scale_e8m0_unbiased = tl.where(scale_e8m0_unbiased > 127, 127, scale_e8m0_unbiased)
# blockscale_e8m0
bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
# Construct quant_scale = 2^(-scale_e8m0_unbiased) via IEEE float bit construction
qs_exp = (-scale_e8m0_unbiased + 127).to(tl.uint32)
quant_scale = (qs_exp << 23).to(tl.float32, bitcast=True)
# Compute quantized x
qx = x * quant_scale
# Convert quantized fp32 tensor to uint32
qx = qx.to(tl.uint32, bitcast=True)
# Extract sign
s = qx & 0x80000000
qx = qx ^ s
qx_fp32 = qx.to(tl.float32, bitcast=True)
saturate_mask = qx_fp32 >= max_normal
denormal_mask = (not saturate_mask) & (qx_fp32 < min_normal)
normal_mask = not (saturate_mask | denormal_mask)
# Denormal numbers
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 numbers
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)
# Merge results
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)
# ========== SplitK fused quant+GEMM kernel for M=16 K=7168 ==========
@triton.jit
def _splitk_fused_quant_gemm_kernel(
a_ptr,
stride_am,
stride_ak,
b_ptr,
stride_bn,
stride_bk,
b_scale_ptr,
stride_bsn,
stride_bsk,
workspace_ptr, # [SPLIT_K, M, N] f32 partial results
stride_ws, # stride for split dimension
stride_wm,
stride_wn,
M,
N,
K,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
BLOCK_SCALE: tl.constexpr,
GROUP_SIZE_M: tl.constexpr,
SPLIT_K: tl.constexpr,
K_PER_SPLIT: tl.constexpr,
):
pid_mn = tl.program_id(0)
pid_k = tl.program_id(1)
num_pid_m = tl.cdiv(M, BLOCK_M)
num_pid_n = tl.cdiv(N, BLOCK_N)
num_pid_in_group = GROUP_SIZE_M * num_pid_n
group_id = pid_mn // num_pid_in_group
first_pid_m = group_id * GROUP_SIZE_M
group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
pid_m = first_pid_m + ((pid_mn % num_pid_in_group) % group_size_m)
pid_n = (pid_mn % num_pid_in_group) // group_size_m
# K range for this split
k_start = pid_k * K_PER_SPLIT
k_end = min(k_start + K_PER_SPLIT, K)
offs_am = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_ak = k_start + tl.arange(0, BLOCK_K)
a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_ak[None, :] * stride_ak
# B pointers start at k_start offset
offs_bn = pid_n * (BLOCK_N // 16) + tl.arange(0, BLOCK_N // 16)
offs_bk_shuffle = tl.arange(0, (BLOCK_K // 2) * 16)
b_ptrs = b_ptr + offs_bn[:, None] * stride_bn + offs_bk_shuffle[None, :] * stride_bk
# Advance B to k_start
b_ptrs += (k_start // 2) * 16 * stride_bk
offs_bsn = pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)
offs_bsk = tl.arange(0, BLOCK_K // BLOCK_SCALE * 32)
b_scale_ptrs = (
b_scale_ptr + offs_bsn[:, None] * stride_bsn + offs_bsk[None, :] * stride_bsk
)
b_scale_ptrs += (k_start // BLOCK_SCALE) * 32 * stride_bsk
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
num_k_iters = tl.cdiv(k_end - k_start, BLOCK_K)
for _ in range(0, num_k_iters):
a_mask = (offs_am[:, None] < M) & (offs_ak[None, :] < k_end)
a = tl.load(a_ptrs, mask=a_mask, other=0.0).to(tl.float32)
a_fp4, a_scales = _mxfp4_quant_op(a, BLOCK_K, BLOCK_M, BLOCK_SCALE)
b_raw = tl.load(b_ptrs, cache_modifier=".cg")
b = (
b_raw.reshape(1, BLOCK_N // 16, BLOCK_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_N, BLOCK_K // 2)
.trans(1, 0)
)
b_scale_raw = tl.load(b_scale_ptrs, cache_modifier=".cg")
b_scales = (
b_scale_raw.reshape(
BLOCK_N // 32,
BLOCK_K // BLOCK_SCALE // 8,
4,
16,
2,
2,
1,
)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_N, BLOCK_K // BLOCK_SCALE)
)
acc = tl.dot_scaled(a_fp4, a_scales, "e2m1", b, b_scales, "e2m1", acc, fast_math=True)
a_ptrs += BLOCK_K * stride_ak
offs_ak += BLOCK_K
b_ptrs += (BLOCK_K // 2) * 16 * stride_bk
b_scale_ptrs += BLOCK_K * stride_bsk
# Store partial result to workspace
offs_cm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_cn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
w_ptrs = (workspace_ptr + pid_k * stride_ws
+ offs_cm[:, None] * stride_wm + offs_cn[None, :] * stride_wn)
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
tl.store(w_ptrs, acc, mask=c_mask)
@triton.jit
def _splitk_reduce_kernel(
workspace_ptr,
stride_ws,
stride_wm,
stride_wn,
c_ptr,
stride_cm,
stride_cn,
M,
N,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
SPLIT_K: 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)
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for s in range(SPLIT_K):
w_ptrs = (workspace_ptr + s * stride_ws
+ offs_m[:, None] * stride_wm + offs_n[None, :] * stride_wn)
partial = tl.load(w_ptrs, mask=mask, other=0.0)
acc += partial
c_ptrs = c_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
tl.store(c_ptrs, acc.to(tl.bfloat16), mask=mask)
# ========== FUSED quant+GEMM kernel (high-upside K=512 prototype) ==========
@triton.jit
def _fused_quant_gemm_kernel(
a_ptr,
stride_am,
stride_ak,
b_ptr,
stride_bn,
stride_bk,
b_scale_ptr,
stride_bsn,
stride_bsk,
c_ptr,
stride_cm,
stride_cn,
M,
N,
K,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
BLOCK_SCALE: tl.constexpr,
GROUP_SIZE_M: tl.constexpr,
):
pid = tl.program_id(0)
num_pid_m = tl.cdiv(M, BLOCK_M)
num_pid_n = tl.cdiv(N, BLOCK_N)
num_pid_in_group = GROUP_SIZE_M * num_pid_n
group_id = pid // num_pid_in_group
first_pid_m = group_id * GROUP_SIZE_M
group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
pid_m = first_pid_m + ((pid % num_pid_in_group) % group_size_m)
pid_n = (pid % num_pid_in_group) // group_size_m
offs_am = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_ak = tl.arange(0, BLOCK_K)
a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_ak[None, :] * stride_ak
offs_bn = pid_n * (BLOCK_N // 16) + tl.arange(0, BLOCK_N // 16)
offs_bk_shuffle = tl.arange(0, (BLOCK_K // 2) * 16)
b_ptrs = b_ptr + offs_bn[:, None] * stride_bn + offs_bk_shuffle[None, :] * stride_bk
offs_bsn = pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)
offs_bsk = tl.arange(0, BLOCK_K // BLOCK_SCALE * 32)
b_scale_ptrs = (
b_scale_ptr + offs_bsn[:, None] * stride_bsn + offs_bsk[None, :] * stride_bsk
)
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for _ in range(0, tl.cdiv(K, BLOCK_K)):
a_mask = (offs_am[:, None] < M) & (offs_ak[None, :] < K)
a = tl.load(a_ptrs, mask=a_mask, other=0.0).to(tl.float32)
a_fp4, a_scales = _mxfp4_quant_op(a, BLOCK_K, BLOCK_M, BLOCK_SCALE)
b_raw = tl.load(b_ptrs, cache_modifier=".cg")
b = (
b_raw.reshape(1, BLOCK_N // 16, BLOCK_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_N, BLOCK_K // 2)
.trans(1, 0)
)
b_scale_raw = tl.load(b_scale_ptrs, cache_modifier=".cg")
b_scales = (
b_scale_raw.reshape(
BLOCK_N // 32,
BLOCK_K // BLOCK_SCALE // 8,
4,
16,
2,
2,
1,
)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_N, BLOCK_K // BLOCK_SCALE)
)
acc = tl.dot_scaled(a_fp4, a_scales, "e2m1", b, b_scales, "e2m1", acc, fast_math=True)
a_ptrs += BLOCK_K * stride_ak
offs_ak += BLOCK_K
b_ptrs += (BLOCK_K // 2) * 16 * stride_bk
b_scale_ptrs += BLOCK_K * stride_bsk
offs_cm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_cn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
c_ptrs = c_ptr + offs_cm[:, None] * stride_cm + offs_cn[None, :] * stride_cn
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
tl.store(c_ptrs, acc.to(tl.bfloat16), mask=c_mask)
# ========== FUSED quant+shuffle kernel (best current fallback path) ==========
@triton.jit
def _fused_quant_shuffle_kernel(
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,
PADDED_N_SCALE,
# B prefetch params (used when ENABLE_B_PREFETCH=True)
b_ptr, b_n_int32, dummy_ptr,
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,
SCALING_MODE: tl.constexpr,
ENABLE_B_PREFETCH: tl.constexpr,
N_QUANT_M: tl.constexpr,
B_PREFETCH_BLOCK: tl.constexpr,
):
pid_m = tl.program_id(0)
# B prefetch path: extra blocks beyond quant range read B into L2
if ENABLE_B_PREFETCH:
if pid_m >= N_QUANT_M:
b_pid = (pid_m - N_QUANT_M) * tl.num_programs(1) + tl.program_id(1)
offs = b_pid * B_PREFETCH_BLOCK + tl.arange(0, B_PREFETCH_BLOCK)
mask = offs < b_n_int32
v = tl.load(b_ptr + offs, mask=mask)
tl.store(dummy_ptr + b_pid, tl.sum(v))
return
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
pn_stride_a = PADDED_N_SCALE * 32
for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), 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(
x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE
)
# Store x_fp4
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)
# Store scale at shuffled positions
row = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
col = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
a = row >> 5
b = (row >> 4) & 1
c = row & 0xF
d = col >> 3
e = (col >> 2) & 1
f = col & 3
shuffled_idx = (
a[:, None] * pn_stride_a
+ d[None, :] * 256
+ f[None, :] * 64
+ c[:, None] * 4
+ e[None, :] * 2
+ b[:, None]
)
if EVEN_M_N:
tl.store(bs_shuffled_ptr + shuffled_idx, bs_e8m0)
else:
n_scale = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
bs_mask = (row < M)[:, None] & (col < n_scale)[None, :]
tl.store(bs_shuffled_ptr + shuffled_idx, bs_e8m0, mask=bs_mask)
# ========== Buffers ==========
_FUSED_BUF = {}
_FUSED_GEMM_OUT_BUF = {}
MXFP4_QUANT_BLOCK_SIZE = SCALE_GROUP_SIZE
_B_PREFETCH_DUMMY = None # small dummy buffer for B-prefetch DCE prevention
def _fused_quant_shuffle(x, b_shuffle=None):
M, N = x.shape
n_scale = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
pm = (M + 255) // 256 * 256
pn = (n_scale + 7) // 8 * 8
key = (M, N)
bufs = _FUSED_BUF.get(key)
if bufs is None:
x_fp4 = torch.empty(M, N // 2, dtype=torch.uint8, device=x.device)
bs_shuffled = torch.zeros(pm * pn, dtype=torch.uint8, device=x.device)
_FUSED_BUF[key] = (x_fp4, bs_shuffled)
else:
x_fp4, bs_shuffled = bufs
# Config selection (matches dynamic_mxfp4_quant wrapper exactly)
if M <= 32:
NUM_ITER = 1
BLOCK_SIZE_M = triton.next_power_of_2(M)
BLOCK_SIZE_N = 32
NUM_WARPS = 1
NUM_STAGES = 1
else:
NUM_ITER = 4
BLOCK_SIZE_M = 64
BLOCK_SIZE_N = 64
NUM_WARPS = 4
NUM_STAGES = 2
if N <= 16384:
BLOCK_SIZE_M = 32
BLOCK_SIZE_N = 128
# Override for small N values
if N <= 1024:
NUM_ITER = 1
NUM_STAGES = 1
NUM_WARPS = 4
BLOCK_SIZE_N = min(256, triton.next_power_of_2(N))
BLOCK_SIZE_N = max(32, BLOCK_SIZE_N)
BLOCK_SIZE_M = min(8, triton.next_power_of_2(M))
# Per-shape override for M=16, K=7168 (worst case shape)
if M == 16 and N == 7168:
BLOCK_SIZE_M = 16
BLOCK_SIZE_N = 128
NUM_ITER = 4
NUM_WARPS = 4
NUM_STAGES = 2
# Per-shape override for M=64, K=2048
if M == 64 and N == 2048:
BLOCK_SIZE_M = 32
BLOCK_SIZE_N = 64
NUM_ITER = 1
NUM_WARPS = 4
NUM_STAGES = 1
# Per-shape override for M=256, K=1536
if M == 256 and N == 1536:
BLOCK_SIZE_M = 32
BLOCK_SIZE_N = 64
NUM_ITER = 1
NUM_WARPS = 4
NUM_STAGES = 1
EVEN_M_N = (M % BLOCK_SIZE_M == 0) and (N % (BLOCK_SIZE_N * NUM_ITER) == 0)
n_quant_m = triton.cdiv(M, BLOCK_SIZE_M)
n_quant_n = triton.cdiv(N, BLOCK_SIZE_N * NUM_ITER)
# B prefetch: expand grid with extra blocks that read B_shuffle into L2
global _B_PREFETCH_DUMMY
ENABLE_B_PREFETCH = b_shuffle is not None
B_PREFETCH_BLOCK = 4096 # int32 elements per B-prefetch block (16KB)
if ENABLE_B_PREFETCH:
b_flat = b_shuffle.view(torch.int32).reshape(-1)
b_n_int32 = b_flat.numel()
n_b_blocks = triton.cdiv(b_n_int32, B_PREFETCH_BLOCK)
# Extra rows in dim 0 for B-prefetch blocks
n_b_extra_rows = triton.cdiv(n_b_blocks, n_quant_n)
grid = (n_quant_m + n_b_extra_rows, n_quant_n)
if _B_PREFETCH_DUMMY is None or _B_PREFETCH_DUMMY.numel() < n_b_blocks:
_B_PREFETCH_DUMMY = torch.empty(n_b_blocks, dtype=torch.int32, device=x.device)
b_ptr = b_flat
dummy_ptr = _B_PREFETCH_DUMMY
else:
grid = (n_quant_m, n_quant_n)
b_ptr = x # dummy pointer (unused)
b_n_int32 = 0
dummy_ptr = x # dummy pointer (unused)
_fused_quant_shuffle_kernel[grid](
x, x_fp4, bs_shuffled,
*x.stride(), *x_fp4.stride(),
M, N, pn,
b_ptr, b_n_int32, dummy_ptr,
BLOCK_SIZE_M=BLOCK_SIZE_M, BLOCK_SIZE_N=BLOCK_SIZE_N,
NUM_ITER=NUM_ITER, NUM_STAGES=NUM_STAGES,
MXFP4_QUANT_BLOCK_SIZE=MXFP4_QUANT_BLOCK_SIZE,
EVEN_M_N=EVEN_M_N, SCALING_MODE=0,
ENABLE_B_PREFETCH=ENABLE_B_PREFETCH,
N_QUANT_M=n_quant_m,
B_PREFETCH_BLOCK=B_PREFETCH_BLOCK,
num_warps=NUM_WARPS, waves_per_eu=2, num_stages=1,
)
return x_fp4, bs_shuffled.view(pm, pn)
# Pre-allocated GEMM output buffers: {(M, N): tensor}
_OUT_BUF = {}
_K32 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_SPLITK_WS = {} # workspace for splitK: {(M, N, SPLIT_K): tensor}
_SPLITK_OUT = {} # output for splitK: {(M, N): tensor}
def _run_fused_smallk_quant_gemm(A, B_shuffle, B_scale_sh):
M, K = A.shape
N = B_shuffle.shape[0]
key = (M, N, K)
out = _FUSED_GEMM_OUT_BUF.get(key)
if out is None:
out = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
_FUSED_GEMM_OUT_BUF[key] = out
# Match the physical storage contract used by aiter's preshuffle kernel.
b_phys = B_shuffle.view(torch.uint8).reshape(N // 16, B_shuffle.shape[1] * 16)
b_scale_phys = B_scale_sh.view(torch.uint8).reshape(
B_scale_sh.shape[0] // 32, B_scale_sh.shape[1] * 32
)
block_m = 16 if M < 16 else triton.next_power_of_2(M)
# For M=32 K=512: use BLOCK_M=16 to double CU utilization
if M == 32 and K == 512:
block_m = 16
block_n = 128
# Shape-scoped warps: more warps for tiny M (better latency hiding),
# standard warps for larger M (less register pressure)
nw = 8 if M <= 4 else 4
wpe = 0 if M <= 4 else 2
grid = (triton.cdiv(M, block_m) * triton.cdiv(N, block_n),)
_fused_quant_gemm_kernel[grid](
A,
A.stride(0),
A.stride(1),
b_phys,
b_phys.stride(0),
b_phys.stride(1),
b_scale_phys,
b_scale_phys.stride(0),
b_scale_phys.stride(1),
out,
out.stride(0),
out.stride(1),
M,
N,
K,
BLOCK_M=block_m,
BLOCK_N=block_n,
BLOCK_K=K,
BLOCK_SCALE=SCALE_GROUP_SIZE,
GROUP_SIZE_M=1,
num_warps=nw,
num_stages=2,
waves_per_eu=wpe,
)
return out
def _run_splitk_fused_quant_gemm(A, B_shuffle, B_scale_sh, SPLIT_K=8, BLOCK_M=16, BLOCK_K=512):
"""SplitK fused quant+GEMM for CU-starved shapes."""
M, K = A.shape
N = B_shuffle.shape[0]
# Pre-allocate workspace and output
ws_key = (M, N, SPLIT_K)
ws = _SPLITK_WS.get(ws_key)
if ws is None:
ws = torch.empty((SPLIT_K, M, N), dtype=torch.float32, device=A.device)
_SPLITK_WS[ws_key] = ws
out_key = (M, N)
out = _SPLITK_OUT.get(out_key)
if out is None:
out = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
_SPLITK_OUT[out_key] = out
# B physical layout (preshuffle)
b_phys = B_shuffle.view(torch.uint8).reshape(N // 16, B_shuffle.shape[1] * 16)
b_scale_phys = B_scale_sh.view(torch.uint8).reshape(
B_scale_sh.shape[0] // 32, B_scale_sh.shape[1] * 32
)
BLOCK_N = 128
K_PER_SPLIT = triton.cdiv(K, SPLIT_K)
# Align K_PER_SPLIT to BLOCK_K
K_PER_SPLIT = ((K_PER_SPLIT + BLOCK_K - 1) // BLOCK_K) * BLOCK_K
num_mn_blocks = triton.cdiv(M, BLOCK_M) * triton.cdiv(N, BLOCK_N)
grid = (num_mn_blocks, SPLIT_K)
_splitk_fused_quant_gemm_kernel[grid](
A, A.stride(0), A.stride(1),
b_phys, b_phys.stride(0), b_phys.stride(1),
b_scale_phys, b_scale_phys.stride(0), b_scale_phys.stride(1),
ws, ws.stride(0), ws.stride(1), ws.stride(2),
M, N, K,
BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
BLOCK_SCALE=SCALE_GROUP_SIZE, GROUP_SIZE_M=1,
SPLIT_K=SPLIT_K, K_PER_SPLIT=K_PER_SPLIT,
num_warps=4, num_stages=2, waves_per_eu=2,
)
# Reduce partial results
reduce_grid = (triton.cdiv(M, BLOCK_M), triton.cdiv(N, BLOCK_N))
_splitk_reduce_kernel[reduce_grid](
ws, ws.stride(0), ws.stride(1), ws.stride(2),
out, out.stride(0), out.stride(1),
M, N,
BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, SPLIT_K=SPLIT_K,
num_warps=4, num_stages=1, waves_per_eu=2,
)
return out
def custom_kernel(data: input_t) -> output_t:
A, _, _, B_shuffle, B_scale_sh = data
M, K = A.shape
N = B_shuffle.shape[0]
if K == 512 and M <= 32:
return _run_fused_smallk_quant_gemm(A, B_shuffle, B_scale_sh)
# SplitK fused quant+GEMM for CU-starved shapes
if M == 16 and K == 7168:
# 17 M*N blocks * 14 K-splits = 238 blocks on 256 CUs (~0.93 waves)
return _run_splitk_fused_quant_gemm(A, B_shuffle, B_scale_sh, SPLIT_K=14, BLOCK_M=16, BLOCK_K=512)
# Always quantize A fresh (no cache — LB uses fresh tensors each call)
# For M≥64: fuse B-prefetch into quant kernel (extra grid blocks read B into L2)
b_prefetch = B_shuffle if M >= 64 else None
A_fp4, A_scale_sh = _fused_quant_shuffle(A, b_shuffle=b_prefetch)
A_fp4_v = A_fp4.view(dtypes.fp4x2)
A_scale_v = A_scale_sh.view(dtypes.fp8_e8m0)
# Direct ASM for all shapes -- bypass wrapper overhead (config CSV reads)
key = (M, N)
out = _OUT_BUF.get(key)
if out is None:
padded_m = ((M + 31) // 32) * 32
out = torch.empty((padded_m, N), dtype=torch.bfloat16, device=A.device)
_OUT_BUF[key] = out
# V10: M=64 no splitK (sk=0), M=256 keep splitK (sk=1)
if M == 256:
log2_sk = 1
else:
log2_sk = 0
gemm_a4w4_asm(
A_fp4_v, B_shuffle, A_scale_v, B_scale_sh,
out, _K32, None, 1.0, 0.0, True, log2_sk,
)
return out[:M]
scrolls · 721 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