submission 711914
sean_nobricks · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 518 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-711914?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:777a939c1d26ec5a5cb8bbd8fdf426283415fe6a16a5c84a1604ea459edd0b21
license declaredunknown
license concludedunknown
authorssean_nobricks
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
autotune
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32, 'BLOCK_K': 256}, num_warps=4, num_stages=2),fp4
"""MXFP4 GEMM with shape-specialized dispatch.num-warps = 4
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32, 'BLOCK_K': 256}, num_warps=4, num_stages=2),split-k
conversion for A and split-K accumulation.stages = 2
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32, 'BLOCK_K': 256}, num_warps=4, num_stages=2),tile-k = 32
BLOCK_K = 32tile-m = 16
BLOCK_M = 16Kernel source
submission.py518 lines
"""MXFP4 GEMM with shape-specialized dispatch.
Path 1 (K <= 512): fused Triton kernel that quantizes A in-register and calls
`tl.dot_scaled`.
Path 2 (K > 512, M > 32): standalone Triton quantization for A followed by a
direct AITER FP4 GEMM call on preshuffled B.
Path 3 (K > 512, M <= 32): fused Triton kernel that uses hardware FP4
conversion for A and split-K accumulation.
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
# =============================================================================
# Software MXFP4 quant — used for K <= 512 fused path (Path 1)
# =============================================================================
@triton.jit
def _mxfp4_quant_tile(x, BLOCK_M: tl.constexpr, BLOCK_K: tl.constexpr):
"""Quantize (BLOCK_M, BLOCK_K) fp32 tile to MXFP4 in-register."""
SG: tl.constexpr = 32
NG: tl.constexpr = BLOCK_K // SG
x = x.reshape(BLOCK_M, NG, SG)
amax = tl.max(tl.abs(x), axis=2, keep_dims=True)
amax_i = amax.to(tl.int32, bitcast=True)
amax_i = (amax_i + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax = amax_i.to(tl.float32, bitcast=True)
scale_ub = tl.log2(amax).floor() - 2.0
scale_ub = tl.clamp(scale_ub, min=-127.0, max=127.0)
scales = scale_ub.to(tl.uint8) + 127
qx = x * tl.exp2(-scale_ub)
qx_u = qx.to(tl.uint32, bitcast=True)
sign = qx_u & 0x80000000
qx_u = qx_u ^ sign
qx_f = qx_u.to(tl.float32, bitcast=True)
sat = qx_f >= 6.0
den = (~sat) & (qx_f < 1.0)
nor = ~(sat | den)
den_x = (qx_f + 4194304.0).to(tl.uint32, bitcast=True) - 1249902592
den_x = den_x.to(tl.uint8)
mant_odd = (qx_u >> 22) & 1
nor_x = qx_u + 0xC11FFFFF
nor_x = nor_x + mant_odd
nor_x = (nor_x >> 22).to(tl.uint8)
e2m1 = tl.full([BLOCK_M, NG, SG], 7, dtype=tl.uint8)
e2m1 = tl.where(nor, nor_x, e2m1)
e2m1 = tl.where(den, den_x, e2m1)
e2m1 = e2m1 | (sign >> 28).to(tl.uint8)
e2m1 = tl.reshape(e2m1, [BLOCK_M, NG, SG // 2, 2])
ev, od = tl.split(e2m1)
fp4 = ev | (od << 4)
return fp4.reshape(BLOCK_M, BLOCK_K // 2), scales.reshape(BLOCK_M, NG)
# =============================================================================
# Hardware MXFP4 quant — used for K > 512, M <= 32 fused path (Path 3)
# =============================================================================
@triton.jit
def _mxfp4_quant_tile_hw(x, BLOCK_M: tl.constexpr, BLOCK_K: tl.constexpr):
"""Quantize (BLOCK_M, BLOCK_K) fp32 tile to MXFP4 via v_cvt_scalef32_pk_fp4_f32.
The instruction computes fp4(src / scale), so passing 2^scale_ub produces
the expected block-scaled quantization. The int32 output dtype with a tied
destination operand matches the packed 32-bit register layout expected by
the instruction.
"""
SG: tl.constexpr = 32
NG: tl.constexpr = BLOCK_K // SG
x = x.reshape(BLOCK_M, NG, SG)
amax = tl.max(tl.abs(x), axis=2, keep_dims=True)
amax_i = amax.to(tl.int32, bitcast=True)
amax_i = (amax_i + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax = amax_i.to(tl.float32, bitcast=True)
scale_ub = tl.log2(amax).floor() - 2.0
scale_ub = tl.clamp(scale_ub, min=-127.0, max=127.0)
scales = scale_ub.to(tl.uint8) + 127
hw_scale = tl.exp2(scale_ub)
hw_scale_broadcast = tl.broadcast_to(hw_scale, (BLOCK_M, NG, SG // 2))
x_pairs = x.reshape(BLOCK_M, NG, SG // 2, 2)
x_even, x_odd = tl.split(x_pairs)
old_vdst = tl.zeros((BLOCK_M, NG, SG // 2), dtype=tl.int32)
fp4_i32 = tl.inline_asm_elementwise(
asm="v_cvt_scalef32_pk_fp4_f32 $0, $1, $2, $3",
constraints="=v,v,v,v,0",
args=[x_even, x_odd, hw_scale_broadcast, old_vdst],
dtype=tl.int32,
is_pure=True,
pack=1,
)
fp4 = (fp4_i32 & 0xFF).to(tl.uint8)
return fp4.reshape(BLOCK_M, BLOCK_K // 2), scales.reshape(BLOCK_M, NG)
# =============================================================================
# B preshuffle unshuffle helper
# =============================================================================
@triton.jit
def _unshuffle_b_preshuffle(b_wide, BN_GROUPS: tl.constexpr, BLOCK_K: tl.constexpr, BLOCK_N: tl.constexpr):
"""Unshuffle preshuffle-loaded B from (BN_GROUPS, WIDE_K) to (BK//2, BN)."""
b = b_wide.reshape(BN_GROUPS, BLOCK_K // 64, 2, 16, 16)
b = b.permute(1, 2, 4, 0, 3)
return b.reshape(BLOCK_K // 2, BLOCK_N)
@triton.jit
def _load_b_scales_from_preshuffled_generic(
b_scale_ptr,
stride_bsn, stride_bsk,
pid_n,
scale_k_start,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
):
"""Load natural-layout B scales for BK values with contiguous shuffled storage."""
SG: tl.constexpr = 32
num_scale_k: tl.constexpr = BLOCK_K // SG
b_scale_block_n = pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)
SHUFFLED_SCALE_K: tl.constexpr = num_scale_k * SG
b_scale_k_offs = tl.arange(0, SHUFFLED_SCALE_K)
scale_k_start_shuffled = scale_k_start * SG
b_scale_ptrs = (
b_scale_ptr
+ b_scale_block_n[:, None] * stride_bsn
+ (scale_k_start_shuffled + b_scale_k_offs[None, :]) * stride_bsk
)
return tl.load(b_scale_ptrs).reshape(
BLOCK_N // 32, BLOCK_K // SG // 8, 4, 16, 2, 2, 1,
).permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_N, num_scale_k)
# =============================================================================
# Path 1: Autotuned fused GEMM kernel (K <= 512)
# =============================================================================
_fused_k512_configs = [
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32, 'BLOCK_K': 256}, num_warps=4, num_stages=2),
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32, 'BLOCK_K': 256}, num_warps=4, num_stages=1),
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64, 'BLOCK_K': 256}, num_warps=4, num_stages=2),
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64, 'BLOCK_K': 256}, num_warps=4, num_stages=1),
triton.Config({'BLOCK_M': 32, 'BLOCK_N': 32, 'BLOCK_K': 256}, num_warps=4, num_stages=2),
triton.Config({'BLOCK_M': 32, 'BLOCK_N': 32, 'BLOCK_K': 256}, num_warps=4, num_stages=1),
triton.Config({'BLOCK_M': 32, 'BLOCK_N': 64, 'BLOCK_K': 256}, num_warps=4, num_stages=2),
triton.Config({'BLOCK_M': 32, 'BLOCK_N': 64, 'BLOCK_K': 256}, num_warps=4, num_stages=1),
triton.Config({'BLOCK_M': 32, 'BLOCK_N': 128, 'BLOCK_K': 256}, num_warps=4, num_stages=2),
triton.Config({'BLOCK_M': 32, 'BLOCK_N': 128, 'BLOCK_K': 256}, num_warps=8, num_stages=2),
]
@triton.autotune(configs=_fused_k512_configs, key=['M', 'N', 'K'])
@triton.jit
def mxfp4_gemm_fused_k512_kernel(
a_ptr, b_ptr, c_ptr, b_scale_ptr,
M, N, K,
stride_am, stride_ak,
stride_bk, stride_bn,
stride_cm, stride_cn,
stride_bsn, stride_bsk,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
):
SG: tl.constexpr = 32
BN_GROUPS: tl.constexpr = BLOCK_N // 16
WIDE_K: tl.constexpr = BLOCK_K // 2 * 16
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(stride_bsn > 0)
tl.assume(stride_bsk > 0)
pid_mn = tl.program_id(0)
num_pid_n = tl.cdiv(N, BLOCK_N)
pid_m = pid_mn // num_pid_n
pid_n = pid_mn % num_pid_n
offs_m = (pid_m * BLOCK_M + tl.arange(0, BLOCK_M)) % M
a_offs_k = tl.arange(0, BLOCK_K)
a_ptrs = a_ptr + offs_m[:, None] * stride_am + a_offs_k[None, :] * stride_ak
offs_bn_groups = (pid_n * BN_GROUPS + tl.arange(0, BN_GROUPS)) % (N // 16)
b_wide_offs = tl.arange(0, WIDE_K)
b_ptrs = b_ptr + offs_bn_groups[:, None] * stride_bn + b_wide_offs[None, :] * stride_bk
NUM_SCALE_K: tl.constexpr = BLOCK_K // SG
accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
num_k_iter = tl.cdiv(K, BLOCK_K)
scale_k_iter_start = 0
for _ in range(0, num_k_iter):
a_bf16 = tl.load(a_ptrs)
a_fp4, a_scales = _mxfp4_quant_tile(a_bf16.to(tl.float32), BLOCK_M, BLOCK_K)
b_wide = tl.load(b_ptrs)
b = _unshuffle_b_preshuffle(b_wide, BN_GROUPS, BLOCK_K, BLOCK_N)
b_scales = _load_b_scales_from_preshuffled_generic(
b_scale_ptr, stride_bsn, stride_bsk, pid_n, scale_k_iter_start, BLOCK_N, BLOCK_K,
)
accumulator += tl.dot_scaled(a_fp4, a_scales, "e2m1", b, b_scales, "e2m1")
a_ptrs += BLOCK_K * stride_ak
b_ptrs += WIDE_K * stride_bk
scale_k_iter_start += NUM_SCALE_K
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, accumulator.to(tl.bfloat16), mask=c_mask)
# =============================================================================
# Path 3: Fused hardware-quant GEMM (K > 512, M <= 32) with split-K
# =============================================================================
_fused_klarge_configs = [
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32, 'BLOCK_K': 256, 'SPLIT_K': 8}, num_warps=4, num_stages=2),
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64, 'BLOCK_K': 256, 'SPLIT_K': 8}, num_warps=4, num_stages=2),
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32, 'BLOCK_K': 256, 'SPLIT_K': 4}, num_warps=4, num_stages=2),
triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64, 'BLOCK_K': 256, 'SPLIT_K': 4}, num_warps=4, num_stages=2),
triton.Config({'BLOCK_M': 32, 'BLOCK_N': 32, 'BLOCK_K': 256, 'SPLIT_K': 4}, num_warps=4, num_stages=2),
triton.Config({'BLOCK_M': 32, 'BLOCK_N': 64, 'BLOCK_K': 256, 'SPLIT_K': 4}, num_warps=4, num_stages=2),
]
@triton.autotune(configs=_fused_klarge_configs, key=['M', 'N', 'K'], reset_to_zero=['c_ptr'])
@triton.jit
def mxfp4_gemm_fused_klarge_kernel(
a_ptr, b_ptr, c_ptr, b_scale_ptr,
M, N, K,
stride_am, stride_ak,
stride_bk, stride_bn,
stride_cm, stride_cn,
stride_bsn, stride_bsk,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
SPLIT_K: tl.constexpr,
):
"""Fused hardware quant + GEMM for K > 512, M <= 32.
A single kernel handles quantization and accumulation together, while
`reset_to_zero` makes autotuned split-K accumulation safe.
"""
SG: tl.constexpr = 32
BN_GROUPS: tl.constexpr = BLOCK_N // 16
WIDE_K: tl.constexpr = BLOCK_K // 2 * 16
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(stride_bsn > 0)
tl.assume(stride_bsk > 0)
pid_mn = tl.program_id(0)
pid_k = tl.program_id(1)
num_pid_n = tl.cdiv(N, BLOCK_N)
pid_m = pid_mn // num_pid_n
pid_n = pid_mn % num_pid_n
offs_m = (pid_m * BLOCK_M + tl.arange(0, BLOCK_M)) % M
k_per_split = tl.cdiv(K, SPLIT_K * BLOCK_K) * BLOCK_K
k_start = pid_k * k_per_split
k_end = tl.minimum(k_start + k_per_split, K)
a_offs_k = tl.arange(0, BLOCK_K)
a_ptrs = a_ptr + offs_m[:, None] * stride_am + (k_start + a_offs_k[None, :]) * stride_ak
offs_bn_groups = (pid_n * BN_GROUPS + tl.arange(0, BN_GROUPS)) % (N // 16)
b_wide_offs = tl.arange(0, WIDE_K)
b_ptrs = b_ptr + offs_bn_groups[:, None] * stride_bn + (k_start * 8 + b_wide_offs[None, :]) * stride_bk
NUM_SCALE_K: tl.constexpr = BLOCK_K // SG
accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
num_k_iter = tl.cdiv(k_end - k_start, BLOCK_K)
scale_k_iter_start = k_start // SG
for _ in range(0, num_k_iter):
a_bf16 = tl.load(a_ptrs)
a_fp4, a_scales = _mxfp4_quant_tile_hw(a_bf16.to(tl.float32), BLOCK_M, BLOCK_K)
b_wide = tl.load(b_ptrs)
b = _unshuffle_b_preshuffle(b_wide, BN_GROUPS, BLOCK_K, BLOCK_N)
b_scales = _load_b_scales_from_preshuffled_generic(
b_scale_ptr, stride_bsn, stride_bsk, pid_n, scale_k_iter_start, BLOCK_N, BLOCK_K,
)
accumulator += tl.dot_scaled(a_fp4, a_scales, "e2m1", b, b_scales, "e2m1")
a_ptrs += BLOCK_K * stride_ak
b_ptrs += WIDE_K * stride_bk
scale_k_iter_start += NUM_SCALE_K
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.atomic_add(c_ptrs, accumulator, mask=c_mask, sem="relaxed")
# =============================================================================
# Standalone A quantization kernel (Path 2: K > 512, M > 32)
# =============================================================================
@triton.jit
def _standalone_quant_kernel(
x_ptr, fp4_ptr, scale_shuffled_ptr,
M, K,
stride_xm, stride_xk,
stride_fm, stride_fk,
SCALE_N_PAD,
BLOCK_M: tl.constexpr, BLOCK_K: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_k = tl.program_id(1)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_k = pid_k * BLOCK_K + tl.arange(0, BLOCK_K)
x_ptrs = x_ptr + offs_m[:, None] * stride_xm + offs_k[None, :] * stride_xk
mask = (offs_m[:, None] < M) & (offs_k[None, :] < K)
x = tl.load(x_ptrs, mask=mask, other=0.0).to(tl.float32)
fp4, scales = _mxfp4_quant_tile_hw(x, BLOCK_M, BLOCK_K)
SG: tl.constexpr = 32
NG: tl.constexpr = BLOCK_K // SG
fp4_offs = pid_k * (BLOCK_K // 2) + tl.arange(0, BLOCK_K // 2)
fp4_ptrs = fp4_ptr + offs_m[:, None] * stride_fm + fp4_offs[None, :] * stride_fk
fp4_mask = (offs_m[:, None] < M) & (fp4_offs[None, :] < K // 2)
tl.store(fp4_ptrs, fp4, mask=fp4_mask)
sc_offs = pid_k * NG + tl.arange(0, NG)
sh_m = offs_m[:, None]
sh_n = sc_offs[None, :]
sh_m_block = sh_m // 32
sh_m_rem = sh_m % 32
sh_m_hi = sh_m_rem // 16
sh_m_lo = sh_m_rem % 16
sh_n_block = sh_n // 8
sh_n_rem = sh_n % 8
sh_n_hi = sh_n_rem // 4
sh_n_lo = sh_n_rem % 4
sc_ptrs = scale_shuffled_ptr + (
sh_m_hi
+ sh_n_hi * 2
+ sh_m_lo * 4
+ sh_n_lo * 64
+ sh_n_block * 256
+ sh_m_block * 32 * SCALE_N_PAD
)
sc_mask = (offs_m[:, None] < M) & (sc_offs[None, :] < K // SG)
tl.store(sc_ptrs, scales, mask=sc_mask)
_quant_buffers = {}
_AITER_ASM_KERNEL_NAME_32X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
def _choose_large_m_aiter_asm_kernel_name(_N, _K):
return _AITER_ASM_KERNEL_NAME_32X128
def _run_aiter_large_m_gemm(A, B_shuffle, B_scale_sh):
import aiter
from aiter import dtypes
A_q_bytes, A_scale_shuffled_bytes = _fast_mxfp4_quant(A)
A_q = A_q_bytes.view(dtypes.fp4x2)
A_scale_shuffled = A_scale_shuffled_bytes.view(dtypes.fp8_e8m0)
M, K = A.shape
N = B_shuffle.shape[0]
padded_m = triton.cdiv(M, 32) * 32
out = torch.empty((padded_m, N), dtype=torch.bfloat16, device=A.device)
kernel_name = _choose_large_m_aiter_asm_kernel_name(N, K)
try:
aiter.gemm_a4w4_asm(
A_q,
B_shuffle,
A_scale_shuffled,
B_scale_sh,
out,
kernel_name,
None,
1.0,
0.0,
True,
0,
)
return out[:M]
except Exception:
return aiter.gemm_a4w4(
A_q,
B_shuffle,
A_scale_shuffled,
B_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
def _fast_mxfp4_quant(A):
"""Standalone MXFP4 quant with pre-allocated buffers."""
M, K = A.shape
key = (M, K)
if key not in _quant_buffers:
scale_m_pad = triton.cdiv(M, 256) * 256
scale_n_pad = triton.cdiv(K // 32, 8) * 8
_quant_buffers[key] = (
torch.empty((M, K // 2), dtype=torch.uint8, device=A.device),
torch.full((scale_m_pad, scale_n_pad), 127, dtype=torch.uint8, device=A.device),
)
fp4, scale_shuffled = _quant_buffers[key]
if M < 32:
BLOCK_M = 16
BLOCK_K = 32
num_warps = 1
else:
BLOCK_M = 32
BLOCK_K = 256
num_warps = 4
scale_n_pad = scale_shuffled.shape[1]
grid = (triton.cdiv(M, BLOCK_M), triton.cdiv(K, BLOCK_K))
_standalone_quant_kernel[grid](
A, fp4, scale_shuffled, M, K,
A.stride(0), A.stride(1),
fp4.stride(0), fp4.stride(1),
scale_n_pad,
BLOCK_M=BLOCK_M, BLOCK_K=BLOCK_K,
num_warps=num_warps,
)
return fp4, scale_shuffled
def custom_kernel(data: input_t) -> output_t:
A, _B, B_q, B_shuffle, B_scale_sh = data
M, K = A.shape
N = B_q.shape[0]
B_scale_raw = B_scale_sh.view(torch.uint8)
padded_N_scale = B_scale_raw.shape[0]
padded_K_scale = B_scale_raw.shape[1]
B_scale_shuffled = B_scale_raw.view(padded_N_scale // 32, padded_K_scale * 32)
B_sh_bytes = B_shuffle.view(torch.uint8)
B_sh_wide = B_sh_bytes.reshape(N // 16, (K // 2) * 16)
if K <= 512:
# Path 1: fused software quant + GEMM, SK=1
C = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
grid = lambda META: (
triton.cdiv(M, META['BLOCK_M']) * triton.cdiv(N, META['BLOCK_N']),
)
mxfp4_gemm_fused_k512_kernel[grid](
A, B_sh_wide, C, B_scale_shuffled,
M, N, K,
A.stride(0), A.stride(1),
B_sh_wide.stride(1), B_sh_wide.stride(0),
C.stride(0), C.stride(1),
B_scale_shuffled.stride(0), B_scale_shuffled.stride(1),
)
return C
elif M <= 32:
# Path 3: fused hardware quant + GEMM, SK=4/8, single launch
C = torch.zeros((M, N), dtype=torch.float32, device=A.device)
grid = lambda META: (
triton.cdiv(M, META['BLOCK_M']) * triton.cdiv(N, META['BLOCK_N']),
META['SPLIT_K'],
)
mxfp4_gemm_fused_klarge_kernel[grid](
A, B_sh_wide, C, B_scale_shuffled,
M, N, K,
A.stride(0), A.stride(1),
B_sh_wide.stride(1), B_sh_wide.stride(0),
C.stride(0), C.stride(1),
B_scale_shuffled.stride(0), B_scale_shuffled.stride(1),
)
return C.to(torch.bfloat16)
else:
return _run_aiter_large_m_gemm(A, B_shuffle, B_scale_sh)
scrolls · 518 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 694791.
- """MXFP4 GEMM — custom Triton FP4 GEMM with hybrid quant dispatch.+ """MXFP4 GEMM with shape-specialized dispatch.- Two kernel paths depending on K:- K <= 512: Single fused kernel — loads bf16 A, quantizes to MXFP4 in-register,- then uses tl.dot_scaled for native FP4 MFMA. Eliminates the separate- quantization kernel launch (~5us overhead).- K > 512: Two kernels — standalone MXFP4 quant (with pre-allocated buffers),- then GEMM kernel on pre-quantized fp4 A. The fused approach is slower- here because 4x larger bf16 A loads per K-iteration dominate.+ Path 1 (K <= 512): fused Triton kernel that quantizes A in-register and calls+ `tl.dot_scaled`.- Both paths use BLOCK_M=16 for M<32 (halves wasted MFMA work) and the Triton- CDNA4 tutorial's in-kernel B scale unshuffle pattern for vectorized scale loads.+ Path 2 (K > 512, M > 32): standalone Triton quantization for A followed by a+ direct AITER FP4 GEMM call on preshuffled B.++ Path 3 (K > 512, M <= 32): fused Triton kernel that uses hardware FP4+ conversion for A and split-K accumulation."""+import torchimport tritonimport triton.language as tl⋯ 1 unchanged lines# =============================================================================- # MXFP4 quantization: bf16 -> fp4(e2m1) + e8m0 block scales- # Adapted from AITER's _mxfp4_quant_op. Used by both the fused GEMM kernel- # (in-register) and the standalone quant kernel (global memory).+ # Software MXFP4 quant — used for K <= 512 fused path (Path 1)# =============================================================================@triton.jitdef _mxfp4_quant_tile(x, BLOCK_M: tl.constexpr, BLOCK_K: tl.constexpr):- """Quantize (BLOCK_M, BLOCK_K) fp32 tile to MXFP4 in-register.-- Returns:- fp4: (BLOCK_M, BLOCK_K // 2) uint8 — nibble-packed e2m1 pairs- scales: (BLOCK_M, BLOCK_K // 32) uint8 — e8m0 block scales- """- SG: tl.constexpr = 32 # scale group size: one e8m0 scale per 32 elements+ """Quantize (BLOCK_M, BLOCK_K) fp32 tile to MXFP4 in-register."""+ SG: tl.constexpr = 32NG: tl.constexpr = BLOCK_K // SGx = x.reshape(BLOCK_M, NG, SG)- # E8M0 block scale: max(|x|) per group, rounded up to nearest power of 2.- # The +0x200000 rounds the fp32 mantissa, &0xFF800000 zeros it out (keeps exponent).amax = tl.max(tl.abs(x), axis=2, keep_dims=True)amax_i = amax.to(tl.int32, bitcast=True)amax_i = (amax_i + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000amax = amax_i.to(tl.float32, bitcast=True)- # Unbiased exponent. The -2 accounts for fp4 e2m1 max value being 6.0 = 2^2 * 1.5scale_ub = tl.log2(amax).floor() - 2.0scale_ub = tl.clamp(scale_ub, min=-127.0, max=127.0)- scales = scale_ub.to(tl.uint8) + 127 # biased e8m0+ scales = scale_ub.to(tl.uint8) + 127- # Scale input into fp4 representable range [0, 6]qx = x * tl.exp2(-scale_ub)- # FP32 -> FP4 (e2m1) conversion via IEEE 754 bit manipulationqx_u = qx.to(tl.uint32, bitcast=True)sign = qx_u & 0x80000000- qx_u = qx_u ^ sign # absolute value+ qx_u = qx_u ^ signqx_f = qx_u.to(tl.float32, bitcast=True)- # Three-way branch: saturate (>=6), denormal (<1), normal (1..6)sat = qx_f >= 6.0den = (~sat) & (qx_f < 1.0)nor = ~(sat | den)- # Denormal path: "magic number" trick — adding 2^22 (=4194304.0) places the- # rounded fp4 bits at known positions in the fp32 mantissa- den_x = (qx_f + 4194304.0).to(tl.uint32, bitcast=True) - 1249902592 # 149 << 23+ den_x = (qx_f + 4194304.0).to(tl.uint32, bitcast=True) - 1249902592den_x = den_x.to(tl.uint8)- # Normal path: adjust exponent bias from fp32 to fp4, round-to-nearest-even- mant_odd = (qx_u >> 22) & 1 # mantissa bit for RTNE- nor_x = qx_u + 0xC11FFFFF # bias adjust: ((1-127) << 23) + (1 << 21) - 1- nor_x = nor_x + mant_odd # RTNE correction+ mant_odd = (qx_u >> 22) & 1+ nor_x = qx_u + 0xC11FFFFF+ nor_x = nor_x + mant_oddnor_x = (nor_x >> 22).to(tl.uint8)- # Merge: default to saturated value 0x7 (max fp4 = 6.0)e2m1 = tl.full([BLOCK_M, NG, SG], 7, dtype=tl.uint8)e2m1 = tl.where(nor, nor_x, e2m1)e2m1 = tl.where(den, den_x, e2m1)- e2m1 = e2m1 | (sign >> 28).to(tl.uint8) # restore sign at bit 3+ e2m1 = e2m1 | (sign >> 28).to(tl.uint8)e2m1 = tl.reshape(e2m1, [BLOCK_M, NG, SG // 2, 2])ev, od = tl.split(e2m1)⋯ 3 unchanged lines# =============================================================================- # Fused GEMM kernel (K <= 512 path)+ # Hardware MXFP4 quant — used for K > 512, M <= 32 fused path (Path 3)# =============================================================================@triton.jit- def mxfp4_gemm_fused_quant_kernel(+ def _mxfp4_quant_tile_hw(x, BLOCK_M: tl.constexpr, BLOCK_K: tl.constexpr):+ """Quantize (BLOCK_M, BLOCK_K) fp32 tile to MXFP4 via v_cvt_scalef32_pk_fp4_f32.++ The instruction computes fp4(src / scale), so passing 2^scale_ub produces+ the expected block-scaled quantization. The int32 output dtype with a tied+ destination operand matches the packed 32-bit register layout expected by+ the instruction.+ """+ SG: tl.constexpr = 32+ NG: tl.constexpr = BLOCK_K // SG++ x = x.reshape(BLOCK_M, NG, SG)++ amax = tl.max(tl.abs(x), axis=2, keep_dims=True)+ amax_i = amax.to(tl.int32, bitcast=True)+ amax_i = (amax_i + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000+ amax = amax_i.to(tl.float32, bitcast=True)+ scale_ub = tl.log2(amax).floor() - 2.0+ scale_ub = tl.clamp(scale_ub, min=-127.0, max=127.0)+ scales = scale_ub.to(tl.uint8) + 127++ hw_scale = tl.exp2(scale_ub)+ hw_scale_broadcast = tl.broadcast_to(hw_scale, (BLOCK_M, NG, SG // 2))++ x_pairs = x.reshape(BLOCK_M, NG, SG // 2, 2)+ x_even, x_odd = tl.split(x_pairs)++ old_vdst = tl.zeros((BLOCK_M, NG, SG // 2), dtype=tl.int32)+ fp4_i32 = tl.inline_asm_elementwise(+ asm="v_cvt_scalef32_pk_fp4_f32 $0, $1, $2, $3",+ constraints="=v,v,v,v,0",+ args=[x_even, x_odd, hw_scale_broadcast, old_vdst],+ dtype=tl.int32,+ is_pure=True,+ pack=1,+ )+ fp4 = (fp4_i32 & 0xFF).to(tl.uint8)++ return fp4.reshape(BLOCK_M, BLOCK_K // 2), scales.reshape(BLOCK_M, NG)+++ # =============================================================================+ # B preshuffle unshuffle helper+ # =============================================================================++ @triton.jit+ def _unshuffle_b_preshuffle(b_wide, BN_GROUPS: tl.constexpr, BLOCK_K: tl.constexpr, BLOCK_N: tl.constexpr):+ """Unshuffle preshuffle-loaded B from (BN_GROUPS, WIDE_K) to (BK//2, BN)."""+ b = b_wide.reshape(BN_GROUPS, BLOCK_K // 64, 2, 16, 16)+ b = b.permute(1, 2, 4, 0, 3)+ return b.reshape(BLOCK_K // 2, BLOCK_N)++ @triton.jit+ def _load_b_scales_from_preshuffled_generic(+ b_scale_ptr,+ stride_bsn, stride_bsk,+ pid_n,+ scale_k_start,+ BLOCK_N: tl.constexpr,+ BLOCK_K: tl.constexpr,+ ):+ """Load natural-layout B scales for BK values with contiguous shuffled storage."""+ SG: tl.constexpr = 32+ num_scale_k: tl.constexpr = BLOCK_K // SG++ b_scale_block_n = pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)+ SHUFFLED_SCALE_K: tl.constexpr = num_scale_k * SG+ b_scale_k_offs = tl.arange(0, SHUFFLED_SCALE_K)+ scale_k_start_shuffled = scale_k_start * SG+ b_scale_ptrs = (+ b_scale_ptr+ + b_scale_block_n[:, None] * stride_bsn+ + (scale_k_start_shuffled + b_scale_k_offs[None, :]) * stride_bsk+ )+ return tl.load(b_scale_ptrs).reshape(+ BLOCK_N // 32, BLOCK_K // SG // 8, 4, 16, 2, 2, 1,+ ).permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_N, num_scale_k)+++ # =============================================================================+ # Path 1: Autotuned fused GEMM kernel (K <= 512)+ # =============================================================================++ _fused_k512_configs = [+ triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32, 'BLOCK_K': 256}, num_warps=4, num_stages=2),+ triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32, 'BLOCK_K': 256}, num_warps=4, num_stages=1),+ triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64, 'BLOCK_K': 256}, num_warps=4, num_stages=2),+ triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64, 'BLOCK_K': 256}, num_warps=4, num_stages=1),+ triton.Config({'BLOCK_M': 32, 'BLOCK_N': 32, 'BLOCK_K': 256}, num_warps=4, num_stages=2),+ triton.Config({'BLOCK_M': 32, 'BLOCK_N': 32, 'BLOCK_K': 256}, num_warps=4, num_stages=1),+ triton.Config({'BLOCK_M': 32, 'BLOCK_N': 64, 'BLOCK_K': 256}, num_warps=4, num_stages=2),+ triton.Config({'BLOCK_M': 32, 'BLOCK_N': 64, 'BLOCK_K': 256}, num_warps=4, num_stages=1),+ triton.Config({'BLOCK_M': 32, 'BLOCK_N': 128, 'BLOCK_K': 256}, num_warps=4, num_stages=2),+ triton.Config({'BLOCK_M': 32, 'BLOCK_N': 128, 'BLOCK_K': 256}, num_warps=8, num_stages=2),+ ]+++ @triton.autotune(configs=_fused_k512_configs, key=['M', 'N', 'K'])+ @triton.jit+ def mxfp4_gemm_fused_k512_kernel(a_ptr, b_ptr, c_ptr, b_scale_ptr,M, N, K,stride_am, stride_ak,⋯ 1 unchanged linesstride_cm, stride_cn,stride_bsn, stride_bsk,BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,- SPLIT_K: tl.constexpr,):SG: tl.constexpr = 32+ BN_GROUPS: tl.constexpr = BLOCK_N // 16+ WIDE_K: tl.constexpr = BLOCK_K // 2 * 16- pid_mn = tl.program_id(0)- pid_k = tl.program_id(1)+ 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(stride_bsn > 0)+ tl.assume(stride_bsk > 0)+ pid_mn = tl.program_id(0)num_pid_n = tl.cdiv(N, BLOCK_N)pid_m = pid_mn // num_pid_npid_n = pid_mn % num_pid_noffs_m = (pid_m * BLOCK_M + tl.arange(0, BLOCK_M)) % M- offs_n = (pid_n * BLOCK_N + tl.arange(0, BLOCK_N)) % N- k_per_split = tl.cdiv(K, SPLIT_K * BLOCK_K) * BLOCK_K- k_start = pid_k * k_per_split- k_end = tl.minimum(k_start + k_per_split, K)-a_offs_k = tl.arange(0, BLOCK_K)- a_ptrs = a_ptr + offs_m[:, None] * stride_am + (k_start + a_offs_k[None, :]) * stride_ak+ a_ptrs = a_ptr + offs_m[:, None] * stride_am + a_offs_k[None, :] * stride_ak- b_offs_k = tl.arange(0, BLOCK_K // 2)- b_ptrs = b_ptr + (k_start // 2 + b_offs_k[:, None]) * stride_bk + offs_n[None, :] * stride_bn+ offs_bn_groups = (pid_n * BN_GROUPS + tl.arange(0, BN_GROUPS)) % (N // 16)+ b_wide_offs = tl.arange(0, WIDE_K)+ b_ptrs = b_ptr + offs_bn_groups[:, None] * stride_bn + b_wide_offs[None, :] * stride_bkNUM_SCALE_K: tl.constexpr = BLOCK_K // SG- SHUFFLED_SCALE_K: tl.constexpr = NUM_SCALE_K * SG- b_scale_block_n = pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)- b_scale_k_offs = tl.arange(0, SHUFFLED_SCALE_K)- scale_k_start = (k_start // SG) * SG- b_scale_ptrs = b_scale_ptr + b_scale_block_n[:, None] * stride_bsn + (scale_k_start + b_scale_k_offs[None, :]) * stride_bskaccumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)- num_k_iter = tl.cdiv(k_end - k_start, BLOCK_K)+ num_k_iter = tl.cdiv(K, BLOCK_K)+ scale_k_iter_start = 0for _ in range(0, num_k_iter):a_bf16 = tl.load(a_ptrs)a_fp4, a_scales = _mxfp4_quant_tile(a_bf16.to(tl.float32), BLOCK_M, BLOCK_K)- b = tl.load(b_ptrs)- # B scales are stored in CDNA4 shuffled layout for coalesced loads.- # Unshuffle in-register via reshape/permute (mfma_nonkdim=16 pattern from- # Triton block-scaled matmul tutorial). Compiler detects this and enables- # 4x vectorized scale loads.- b_scales = tl.load(b_scale_ptrs).reshape(- BLOCK_N // 32, NUM_SCALE_K // 8, 4, 16, 2, 2, 1,- ).permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_N, NUM_SCALE_K)+ b_wide = tl.load(b_ptrs)+ b = _unshuffle_b_preshuffle(b_wide, BN_GROUPS, BLOCK_K, BLOCK_N)+ b_scales = _load_b_scales_from_preshuffled_generic(+ b_scale_ptr, stride_bsn, stride_bsk, pid_n, scale_k_iter_start, BLOCK_N, BLOCK_K,+ )+accumulator += tl.dot_scaled(a_fp4, a_scales, "e2m1", b, b_scales, "e2m1")a_ptrs += BLOCK_K * stride_ak- b_ptrs += (BLOCK_K // 2) * stride_bk- b_scale_ptrs += SHUFFLED_SCALE_K * stride_bsk+ b_ptrs += WIDE_K * stride_bk+ scale_k_iter_start += NUM_SCALE_Koffs_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_cnc_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)+ tl.store(c_ptrs, accumulator.to(tl.bfloat16), mask=c_mask)- if SPLIT_K == 1:- tl.store(c_ptrs, accumulator.to(tl.bfloat16), mask=c_mask)- else:- tl.atomic_add(c_ptrs, accumulator, mask=c_mask, sem="relaxed")-# =============================================================================- # Pre-quantized GEMM kernel (K > 512 path)+ # Path 3: Fused hardware-quant GEMM (K > 512, M <= 32) with split-K# =============================================================================+ _fused_klarge_configs = [+ triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32, 'BLOCK_K': 256, 'SPLIT_K': 8}, num_warps=4, num_stages=2),+ triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64, 'BLOCK_K': 256, 'SPLIT_K': 8}, num_warps=4, num_stages=2),+ triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32, 'BLOCK_K': 256, 'SPLIT_K': 4}, num_warps=4, num_stages=2),+ triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64, 'BLOCK_K': 256, 'SPLIT_K': 4}, num_warps=4, num_stages=2),+ triton.Config({'BLOCK_M': 32, 'BLOCK_N': 32, 'BLOCK_K': 256, 'SPLIT_K': 4}, num_warps=4, num_stages=2),+ triton.Config({'BLOCK_M': 32, 'BLOCK_N': 64, 'BLOCK_K': 256, 'SPLIT_K': 4}, num_warps=4, num_stages=2),+ ]+++ @triton.autotune(configs=_fused_klarge_configs, key=['M', 'N', 'K'], reset_to_zero=['c_ptr'])@triton.jit- def mxfp4_gemm_splitk_kernel(- a_ptr, b_ptr, c_ptr, a_scale_ptr, b_scale_ptr,- M, N, K_packed,+ def mxfp4_gemm_fused_klarge_kernel(+ a_ptr, b_ptr, c_ptr, b_scale_ptr,+ M, N, K,stride_am, stride_ak,stride_bk, stride_bn,stride_cm, stride_cn,- stride_asm, stride_ask,stride_bsn, stride_bsk,BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,SPLIT_K: tl.constexpr,):+ """Fused hardware quant + GEMM for K > 512, M <= 32.++ A single kernel handles quantization and accumulation together, while+ `reset_to_zero` makes autotuned split-K accumulation safe.+ """SG: tl.constexpr = 32+ BN_GROUPS: tl.constexpr = BLOCK_N // 16+ WIDE_K: tl.constexpr = BLOCK_K // 2 * 16+ 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(stride_bsn > 0)+ tl.assume(stride_bsk > 0)+pid_mn = tl.program_id(0)pid_k = tl.program_id(1)⋯ 2 unchanged linespid_n = pid_mn % num_pid_noffs_m = (pid_m * BLOCK_M + tl.arange(0, BLOCK_M)) % M- offs_n = (pid_n * BLOCK_N + tl.arange(0, BLOCK_N)) % N- k_per_split = tl.cdiv(K_packed, SPLIT_K * (BLOCK_K // 2)) * (BLOCK_K // 2)+ k_per_split = tl.cdiv(K, SPLIT_K * BLOCK_K) * BLOCK_Kk_start = pid_k * k_per_split- k_end = tl.minimum(k_start + k_per_split, K_packed)+ k_end = tl.minimum(k_start + k_per_split, K)- offs_k = tl.arange(0, BLOCK_K // 2)+ a_offs_k = tl.arange(0, BLOCK_K)+ a_ptrs = a_ptr + offs_m[:, None] * stride_am + (k_start + a_offs_k[None, :]) * stride_ak- a_ptrs = a_ptr + offs_m[:, None] * stride_am + (k_start + offs_k[None, :]) * stride_ak- b_ptrs = b_ptr + (k_start + offs_k[:, None]) * stride_bk + offs_n[None, :] * stride_bn+ offs_bn_groups = (pid_n * BN_GROUPS + tl.arange(0, BN_GROUPS)) % (N // 16)+ b_wide_offs = tl.arange(0, WIDE_K)+ b_ptrs = b_ptr + offs_bn_groups[:, None] * stride_bn + (k_start * 8 + b_wide_offs[None, :]) * stride_bk- num_scale_k: tl.constexpr = BLOCK_K // SG- offs_sk = tl.arange(0, num_scale_k)- scale_k_start = k_start * 2 // SG- a_scale_ptrs = a_scale_ptr + offs_m[:, None] * stride_asm + (scale_k_start + offs_sk[None, :]) * stride_ask+ NUM_SCALE_K: tl.constexpr = BLOCK_K // SG- SHUFFLED_SCALE_K: tl.constexpr = BLOCK_K // SG * SG- b_scale_block_n = pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)- b_scale_k_offs = tl.arange(0, SHUFFLED_SCALE_K)- scale_k_start_shuffled = scale_k_start * SG- b_scale_ptrs = b_scale_ptr + b_scale_block_n[:, None] * stride_bsn + (scale_k_start_shuffled + b_scale_k_offs[None, :]) * stride_bsk-accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)- num_k_iter = tl.cdiv(k_end - k_start, BLOCK_K // 2)+ num_k_iter = tl.cdiv(k_end - k_start, BLOCK_K)+ scale_k_iter_start = k_start // SGfor _ in range(0, num_k_iter):- a = tl.load(a_ptrs)- b = tl.load(b_ptrs)- a_scales = tl.load(a_scale_ptrs)+ a_bf16 = tl.load(a_ptrs)+ a_fp4, a_scales = _mxfp4_quant_tile_hw(a_bf16.to(tl.float32), BLOCK_M, BLOCK_K)- b_scales = tl.load(b_scale_ptrs).reshape(- BLOCK_N // 32, BLOCK_K // SG // 8, 4, 16, 2, 2, 1,- ).permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_N, BLOCK_K // SG)+ b_wide = tl.load(b_ptrs)+ b = _unshuffle_b_preshuffle(b_wide, BN_GROUPS, BLOCK_K, BLOCK_N)- accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")+ b_scales = _load_b_scales_from_preshuffled_generic(+ b_scale_ptr, stride_bsn, stride_bsk, pid_n, scale_k_iter_start, BLOCK_N, BLOCK_K,+ )- a_ptrs += (BLOCK_K // 2) * stride_ak- b_ptrs += (BLOCK_K // 2) * stride_bk- a_scale_ptrs += num_scale_k * stride_ask- b_scale_ptrs += SHUFFLED_SCALE_K * stride_bsk+ accumulator += tl.dot_scaled(a_fp4, a_scales, "e2m1", b, b_scales, "e2m1")+ a_ptrs += BLOCK_K * stride_ak+ b_ptrs += WIDE_K * stride_bk+ scale_k_iter_start += NUM_SCALE_K+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_cnc_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)+ tl.atomic_add(c_ptrs, accumulator, mask=c_mask, sem="relaxed")- if SPLIT_K == 1:- tl.store(c_ptrs, accumulator.to(tl.bfloat16), mask=c_mask)- else:- tl.atomic_add(c_ptrs, accumulator, mask=c_mask, sem="relaxed")-# =============================================================================- # Standalone A quantization kernel (replaces AITER's dynamic_mxfp4_quant)- # Pre-allocates output buffers per shape to avoid tensor allocation overhead.+ # Standalone A quantization kernel (Path 2: K > 512, M > 32)# =============================================================================@triton.jitdef _standalone_quant_kernel(- x_ptr, fp4_ptr, scale_ptr,+ x_ptr, fp4_ptr, scale_shuffled_ptr,M, K,stride_xm, stride_xk,stride_fm, stride_fk,- stride_sm, stride_sk,+ SCALE_N_PAD,BLOCK_M: tl.constexpr, BLOCK_K: tl.constexpr,):pid_m = tl.program_id(0)⋯ 6 unchanged linesmask = (offs_m[:, None] < M) & (offs_k[None, :] < K)x = tl.load(x_ptrs, mask=mask, other=0.0).to(tl.float32)- fp4, scales = _mxfp4_quant_tile(x, BLOCK_M, BLOCK_K)+ fp4, scales = _mxfp4_quant_tile_hw(x, BLOCK_M, BLOCK_K)SG: tl.constexpr = 32NG: tl.constexpr = BLOCK_K // SG⋯ 4 unchanged linestl.store(fp4_ptrs, fp4, mask=fp4_mask)sc_offs = pid_k * NG + tl.arange(0, NG)- sc_ptrs = scale_ptr + offs_m[:, None] * stride_sm + sc_offs[None, :] * stride_sk+ sh_m = offs_m[:, None]+ sh_n = sc_offs[None, :]+ sh_m_block = sh_m // 32+ sh_m_rem = sh_m % 32+ sh_m_hi = sh_m_rem // 16+ sh_m_lo = sh_m_rem % 16+ sh_n_block = sh_n // 8+ sh_n_rem = sh_n % 8+ sh_n_hi = sh_n_rem // 4+ sh_n_lo = sh_n_rem % 4+ sc_ptrs = scale_shuffled_ptr + (+ sh_m_hi+ + sh_n_hi * 2+ + sh_m_lo * 4+ + sh_n_lo * 64+ + sh_n_block * 256+ + sh_m_block * 32 * SCALE_N_PAD+ )sc_mask = (offs_m[:, None] < M) & (sc_offs[None, :] < K // SG)tl.store(sc_ptrs, scales, mask=sc_mask)_quant_buffers = {}+ _AITER_ASM_KERNEL_NAME_32X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"+ def _choose_large_m_aiter_asm_kernel_name(_N, _K):+ return _AITER_ASM_KERNEL_NAME_32X128++ def _run_aiter_large_m_gemm(A, B_shuffle, B_scale_sh):+ import aiter+ from aiter import dtypes++ A_q_bytes, A_scale_shuffled_bytes = _fast_mxfp4_quant(A)+ A_q = A_q_bytes.view(dtypes.fp4x2)+ A_scale_shuffled = A_scale_shuffled_bytes.view(dtypes.fp8_e8m0)++ M, K = A.shape+ N = B_shuffle.shape[0]+ padded_m = triton.cdiv(M, 32) * 32+ out = torch.empty((padded_m, N), dtype=torch.bfloat16, device=A.device)+ kernel_name = _choose_large_m_aiter_asm_kernel_name(N, K)++ try:+ aiter.gemm_a4w4_asm(+ A_q,+ B_shuffle,+ A_scale_shuffled,+ B_scale_sh,+ out,+ kernel_name,+ None,+ 1.0,+ 0.0,+ True,+ 0,+ )+ return out[:M]+ except Exception:+ return aiter.gemm_a4w4(+ A_q,+ B_shuffle,+ A_scale_shuffled,+ B_scale_sh,+ dtype=dtypes.bf16,+ bpreshuffle=True,+ )++def _fast_mxfp4_quant(A):- """Standalone MXFP4 quant with pre-allocated buffers. Bypasses AITER overhead."""+ """Standalone MXFP4 quant with pre-allocated buffers."""M, K = A.shapekey = (M, K)if key not in _quant_buffers:+ scale_m_pad = triton.cdiv(M, 256) * 256+ scale_n_pad = triton.cdiv(K // 32, 8) * 8_quant_buffers[key] = (torch.empty((M, K // 2), dtype=torch.uint8, device=A.device),- torch.empty((M, K // 32), dtype=torch.uint8, device=A.device),+ torch.full((scale_m_pad, scale_n_pad), 127, dtype=torch.uint8, device=A.device),)- fp4, scale = _quant_buffers[key]- BLOCK_M = 16 if M < 32 else 32- grid = (triton.cdiv(M, BLOCK_M), triton.cdiv(K, 256))+ fp4, scale_shuffled = _quant_buffers[key]+ if M < 32:+ BLOCK_M = 16+ BLOCK_K = 32+ num_warps = 1+ else:+ BLOCK_M = 32+ BLOCK_K = 256+ num_warps = 4+ scale_n_pad = scale_shuffled.shape[1]+ grid = (triton.cdiv(M, BLOCK_M), triton.cdiv(K, BLOCK_K))_standalone_quant_kernel[grid](- A, fp4, scale, M, K,+ A, fp4, scale_shuffled, M, K,A.stride(0), A.stride(1),fp4.stride(0), fp4.stride(1),- scale.stride(0), scale.stride(1),- BLOCK_M=BLOCK_M, BLOCK_K=256,+ scale_n_pad,+ BLOCK_M=BLOCK_M, BLOCK_K=BLOCK_K,+ num_warps=num_warps,)- return fp4, scale+ return fp4, scale_shuffled- # =============================================================================- # Dispatch logic- # =============================================================================-- def _choose_tile_config(M):- """Per-shape tile selection. BLOCK_M=16 for small M reduces MFMA waste."""- BLOCK_K = 256- BLOCK_N = 32- BLOCK_M = 16 if M < 32 else 32- return BLOCK_M, BLOCK_N, BLOCK_K--- def _choose_split_k(M, K, block_k=256):- if M > 32:- return 1- max_useful = K // block_k- if max_useful <= 2:- return 1- return min(8, max_useful)--def custom_kernel(data: input_t) -> output_t:- A, B, B_q, B_shuffle, B_scale_sh = data+ A, _B, B_q, B_shuffle, B_scale_sh = dataM, K = A.shapeN = B_q.shape[0]⋯ 1 unchanged linespadded_N_scale = B_scale_raw.shape[0]padded_K_scale = B_scale_raw.shape[1]B_scale_shuffled = B_scale_raw.view(padded_N_scale // 32, padded_K_scale * 32)- B_q_bytes = B_q.view(torch.uint8)- BLOCK_M, BLOCK_N, BLOCK_K = _choose_tile_config(M)- SPLIT_K = _choose_split_k(M, K, BLOCK_K)+ B_sh_bytes = B_shuffle.view(torch.uint8)+ B_sh_wide = B_sh_bytes.reshape(N // 16, (K // 2) * 16)- out_dtype = torch.float32 if SPLIT_K > 1 else torch.bfloat16- C = torch.zeros((M, N), dtype=out_dtype, device=A.device) if SPLIT_K > 1 else torch.empty((M, N), dtype=out_dtype, device=A.device)- grid = (triton.cdiv(M, BLOCK_M) * triton.cdiv(N, BLOCK_N), SPLIT_K)-if K <= 512:- mxfp4_gemm_fused_quant_kernel[grid](- A, B_q_bytes, C, B_scale_shuffled,+ # Path 1: fused software quant + GEMM, SK=1+ C = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)+ grid = lambda META: (+ triton.cdiv(M, META['BLOCK_M']) * triton.cdiv(N, META['BLOCK_N']),+ )+ mxfp4_gemm_fused_k512_kernel[grid](+ A, B_sh_wide, C, B_scale_shuffled,M, N, K,A.stride(0), A.stride(1),- B_q_bytes.stride(1), B_q_bytes.stride(0),+ B_sh_wide.stride(1), B_sh_wide.stride(0),C.stride(0), C.stride(1),B_scale_shuffled.stride(0), B_scale_shuffled.stride(1),- BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,- SPLIT_K=SPLIT_K,- num_warps=4, num_stages=2,)- else:- A_q, A_scale = _fast_mxfp4_quant(A)- K_packed = K // 2+ return C- mxfp4_gemm_splitk_kernel[grid](- A_q, B_q_bytes, C, A_scale, B_scale_shuffled,- M, N, K_packed,- A_q.stride(0), A_q.stride(1),- B_q_bytes.stride(1), B_q_bytes.stride(0),+ elif M <= 32:+ # Path 3: fused hardware quant + GEMM, SK=4/8, single launch+ C = torch.zeros((M, N), dtype=torch.float32, device=A.device)+ grid = lambda META: (+ triton.cdiv(M, META['BLOCK_M']) * triton.cdiv(N, META['BLOCK_N']),+ META['SPLIT_K'],+ )+ mxfp4_gemm_fused_klarge_kernel[grid](+ A, B_sh_wide, C, B_scale_shuffled,+ M, N, K,+ A.stride(0), A.stride(1),+ B_sh_wide.stride(1), B_sh_wide.stride(0),C.stride(0), C.stride(1),- A_scale.stride(0), A_scale.stride(1),B_scale_shuffled.stride(0), B_scale_shuffled.stride(1),- BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,- SPLIT_K=SPLIT_K,- num_warps=4, num_stages=2,)+ return C.to(torch.bfloat16)- if SPLIT_K > 1:- C = C.to(torch.bfloat16)- return C+ else:+ return _run_aiter_large_m_gemm(A, B_shuffle, B_scale_sh)
scrolls · 680 diff lines total
Best evidence level for this revision: reported
JSON