submission 694791
sean_nobricks · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 379 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-694791?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:5588a86e0ed1e7a3ae3a1ca48aa4e8aa93cb2fad1029bc71c500c6fb49abfcb7
license declaredunknown
license concludedunknown
authorssean_nobricks
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
"""MXFP4 GEMM — custom Triton FP4 GEMM with hybrid quant dispatch.num-warps = 4
num_warps=4, num_stages=2,split-k
SPLIT_K: tl.constexpr,stages = 2
num_warps=4, num_stages=2,tile-k = 256
BLOCK_M=BLOCK_M, BLOCK_K=256,tile-m = 16
Both paths use BLOCK_M=16 for M<32 (halves wasted MFMA work) and the Tritontile-n = 32
BLOCK_N = 32Kernel source
submission.py379 lines
"""MXFP4 GEMM — custom Triton FP4 GEMM with hybrid quant 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.
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.
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
# =============================================================================
# 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).
# =============================================================================
@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.
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
NG: tl.constexpr = BLOCK_K // SG
x = 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) & 0xFF800000
amax = amax_i.to(tl.float32, bitcast=True)
# Unbiased exponent. The -2 accounts for fp4 e2m1 max value being 6.0 = 2^2 * 1.5
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 # biased e8m0
# Scale input into fp4 representable range [0, 6]
qx = x * tl.exp2(-scale_ub)
# FP32 -> FP4 (e2m1) conversion via IEEE 754 bit manipulation
qx_u = qx.to(tl.uint32, bitcast=True)
sign = qx_u & 0x80000000
qx_u = qx_u ^ sign # absolute value
qx_f = qx_u.to(tl.float32, bitcast=True)
# Three-way branch: saturate (>=6), denormal (<1), normal (1..6)
sat = qx_f >= 6.0
den = (~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 = 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
nor_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 = 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)
# =============================================================================
# Fused GEMM kernel (K <= 512 path)
# =============================================================================
@triton.jit
def mxfp4_gemm_fused_quant_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,
):
SG: tl.constexpr = 32
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
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
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
NUM_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_bsk
accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
num_k_iter = tl.cdiv(k_end - k_start, BLOCK_K)
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 = 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)
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
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)
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)
# =============================================================================
@triton.jit
def mxfp4_gemm_splitk_kernel(
a_ptr, b_ptr, c_ptr, a_scale_ptr, b_scale_ptr,
M, N, K_packed,
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,
):
SG: tl.constexpr = 32
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
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_start = pid_k * k_per_split
k_end = tl.minimum(k_start + k_per_split, K_packed)
offs_k = tl.arange(0, BLOCK_K // 2)
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
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
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)
for _ in range(0, num_k_iter):
a = tl.load(a_ptrs)
b = tl.load(b_ptrs)
a_scales = tl.load(a_scale_ptrs)
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)
accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")
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
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)
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.
# =============================================================================
@triton.jit
def _standalone_quant_kernel(
x_ptr, fp4_ptr, scale_ptr,
M, K,
stride_xm, stride_xk,
stride_fm, stride_fk,
stride_sm, stride_sk,
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(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)
sc_ptrs = scale_ptr + offs_m[:, None] * stride_sm + sc_offs[None, :] * stride_sk
sc_mask = (offs_m[:, None] < M) & (sc_offs[None, :] < K // SG)
tl.store(sc_ptrs, scales, mask=sc_mask)
_quant_buffers = {}
def _fast_mxfp4_quant(A):
"""Standalone MXFP4 quant with pre-allocated buffers. Bypasses AITER overhead."""
M, K = A.shape
key = (M, K)
if key not in _quant_buffers:
_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),
)
fp4, scale = _quant_buffers[key]
BLOCK_M = 16 if M < 32 else 32
grid = (triton.cdiv(M, BLOCK_M), triton.cdiv(K, 256))
_standalone_quant_kernel[grid](
A, fp4, scale, 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,
)
return fp4, scale
# =============================================================================
# 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
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_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)
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,
M, N, K,
A.stride(0), A.stride(1),
B_q_bytes.stride(1), B_q_bytes.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
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),
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,
)
if SPLIT_K > 1:
C = C.to(torch.bfloat16)
return C
scrolls · 379 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 690748.
- """MXFP4 GEMM — custom Triton kernel with in-kernel scale unshuffle and split-K."""+ """MXFP4 GEMM — custom Triton FP4 GEMM with hybrid quant 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.++ 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.+ """import torchimport tritonimport triton.language as tlfrom task import input_t, output_t- SCALE_GROUP_SIZE = 32+ # =============================================================================+ # 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).+ # =============================================================================@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.++ 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+ NG: tl.constexpr = BLOCK_K // SG++ x = 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) & 0xFF800000+ amax = amax_i.to(tl.float32, bitcast=True)+ # Unbiased exponent. The -2 accounts for fp4 e2m1 max value being 6.0 = 2^2 * 1.5+ 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 # biased e8m0++ # Scale input into fp4 representable range [0, 6]+ qx = x * tl.exp2(-scale_ub)++ # FP32 -> FP4 (e2m1) conversion via IEEE 754 bit manipulation+ qx_u = qx.to(tl.uint32, bitcast=True)+ sign = qx_u & 0x80000000+ qx_u = qx_u ^ sign # absolute value+ qx_f = qx_u.to(tl.float32, bitcast=True)++ # Three-way branch: saturate (>=6), denormal (<1), normal (1..6)+ sat = qx_f >= 6.0+ den = (~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 = 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+ nor_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 = 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)+++ # =============================================================================+ # Fused GEMM kernel (K <= 512 path)+ # =============================================================================++ @triton.jit+ def mxfp4_gemm_fused_quant_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,+ ):+ SG: tl.constexpr = 32++ 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+ 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++ 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++ NUM_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_bsk++ accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)++ num_k_iter = tl.cdiv(k_end - k_start, BLOCK_K)+ 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 = 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)++ 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++ 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)++ 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)+ # =============================================================================++ @triton.jitdef mxfp4_gemm_splitk_kernel(a_ptr, b_ptr, c_ptr, a_scale_ptr, b_scale_ptr,M, N, K_packed,⋯ 26 unchanged linesa_ptrs = a_ptr + offs_m[:, None] * stride_am + (k_start + offs_k[None, :]) * stride_akb_ptrs = b_ptr + (k_start + offs_k[:, None]) * stride_bk + offs_n[None, :] * stride_bn- # A scales: un-shuffled natural (M, K_scale) layoutnum_scale_k: tl.constexpr = BLOCK_K // SGoffs_sk = tl.arange(0, num_scale_k)scale_k_start = k_start * 2 // SGa_scale_ptrs = a_scale_ptr + offs_m[:, None] * stride_asm + (scale_k_start + offs_sk[None, :]) * stride_ask- # B scales: shuffled (N//32, K_scale*32) layout per Triton CDNA4 tutorial.- # Load contiguous shuffled block, reshape/permute in-register to recover- # logical (BLOCK_N, BLOCK_K//SG) layout. Compiler detects this pattern- # and enables 4x vectorized scale loads.SHUFFLED_SCALE_K: tl.constexpr = BLOCK_K // SG * SGb_scale_block_n = pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)b_scale_k_offs = tl.arange(0, SHUFFLED_SCALE_K)⋯ 8 unchanged linesb = tl.load(b_ptrs)a_scales = tl.load(a_scale_ptrs)- # B scales: load shuffled, unshuffle in-register (mfma_nonkdim=16 pattern)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)⋯ 16 unchanged linestl.atomic_add(c_ptrs, accumulator, mask=c_mask, sem="relaxed")- def _choose_split_k(M, K_packed, block_k_half=128):+ # =============================================================================+ # Standalone A quantization kernel (replaces AITER's dynamic_mxfp4_quant)+ # Pre-allocates output buffers per shape to avoid tensor allocation overhead.+ # =============================================================================++ @triton.jit+ def _standalone_quant_kernel(+ x_ptr, fp4_ptr, scale_ptr,+ M, K,+ stride_xm, stride_xk,+ stride_fm, stride_fk,+ stride_sm, stride_sk,+ 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(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)+ sc_ptrs = scale_ptr + offs_m[:, None] * stride_sm + sc_offs[None, :] * stride_sk+ sc_mask = (offs_m[:, None] < M) & (sc_offs[None, :] < K // SG)+ tl.store(sc_ptrs, scales, mask=sc_mask)+++ _quant_buffers = {}+++ def _fast_mxfp4_quant(A):+ """Standalone MXFP4 quant with pre-allocated buffers. Bypasses AITER overhead."""+ M, K = A.shape+ key = (M, K)+ if key not in _quant_buffers:+ _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),+ )+ fp4, scale = _quant_buffers[key]+ BLOCK_M = 16 if M < 32 else 32+ grid = (triton.cdiv(M, BLOCK_M), triton.cdiv(K, 256))+ _standalone_quant_kernel[grid](+ A, fp4, scale, 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,+ )+ return fp4, scale+++ # =============================================================================+ # 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_packed // block_k_half+ max_useful = K // block_kif max_useful <= 2:return 1return min(8, max_useful)⋯ 3 unchanged linesA, B, B_q, B_shuffle, B_scale_sh = dataM, K = A.shapeN = B_q.shape[0]- K_packed = K // 2- K_scale = K // SCALE_GROUP_SIZE- from aiter.ops.triton.quant import dynamic_mxfp4_quant-- A_fp4, A_scale_raw = dynamic_mxfp4_quant(A)- A_q = A_fp4.view(torch.uint8)- A_scale = A_scale_raw.view(torch.uint8)-- # B scales: reshape AITER's shuffled layout to (N//32, K_scale*32) for- # in-kernel unshuffle. Free view, no data copy.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_q_bytes = B_q.view(torch.uint8)- BLOCK_M, BLOCK_N, BLOCK_K = 32, 32, 256- SPLIT_K = _choose_split_k(M, K_packed)+ BLOCK_M, BLOCK_N, BLOCK_K = _choose_tile_config(M)+ SPLIT_K = _choose_split_k(M, K, BLOCK_K)out_dtype = torch.float32 if SPLIT_K > 1 else torch.bfloat16- if SPLIT_K > 1:- C = torch.zeros((M, N), dtype=out_dtype, device=A.device)+ 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,+ M, N, K,+ A.stride(0), A.stride(1),+ B_q_bytes.stride(1), B_q_bytes.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:- C = torch.empty((M, N), dtype=out_dtype, device=A.device)+ A_q, A_scale = _fast_mxfp4_quant(A)+ K_packed = K // 2- num_m_tiles = triton.cdiv(M, BLOCK_M)- num_n_tiles = triton.cdiv(N, BLOCK_N)- grid = (num_m_tiles * num_n_tiles, SPLIT_K)+ 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),+ 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,+ )- 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),- 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,- )-if SPLIT_K > 1:C = C.to(torch.bfloat16)return C
scrolls · 371 diff lines total
Best evidence level for this revision: reported
JSON