submission 636064
Divyansh Khanna · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 851 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-636064?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:b6aa5d342367519994484d0d290cccfe25339a6e1261be42ebd5be18269bd5a3
license declaredunknown
license concludedunknown
authorsDivyansh Khanna
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.num-warps = 1
NUM_WARPS = 1split-k
This version (1 kernel launch + optional splitK reduce):stages = 1
NUM_STAGES = 1tile-m = 64
BLOCK_SIZE_M = 64tile-n = 32
BLOCK_SIZE_N = 32Kernel source
submission.py851 lines
"""
FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.
Quant logic follows aiter op_tests/test_gemm_a4w4.py (get_triton_quant(QuantType.per_1x32)).
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
# Reuse aiter's quantization math — no point duplicating 100 lines of FP4 bit manipulation
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
# ─────────────────────────────────────────────────────────────────────────────
# Patched Triton kernel: dynamic_mxfp4_quant with inline e8m0 scale shuffle.
#
# This is aiter's _dynamic_mxfp4_quant_kernel with one addition: the scale
# store section has a SHUFFLE branch that writes scales in the permuted layout
# gemm_a4w4 expects. The shuffle index math is copied verbatim from aiter's
# _fused_rms_mxfp4_quant_kernel (fused_mxfp4_quant.py:173-191).
# ─────────────────────────────────────────────────────────────────────────────
@triton.heuristics(
{
"EVEN_M_N": lambda args: args["M"] % args["BLOCK_SIZE_M"] == 0
and args["N"] % (args["BLOCK_SIZE_N"] * args["NUM_ITER"]) == 0,
}
)
@triton.jit
def _dynamic_mxfp4_quant_shuffle_kernel(
# Pointers
x_ptr,
x_fp4_ptr,
bs_ptr,
# Strides
stride_x_m_in,
stride_x_n_in,
stride_x_fp4_m_in,
stride_x_fp4_n_in,
stride_bs_m_in,
stride_bs_n_in,
# Problem size
M,
N,
# Compile-time constants
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,
# --- Shuffle support (new vs aiter's original) ---
SHUFFLE: tl.constexpr, # Whether to write scales in shuffled order
SCALE_N_PAD: tl.constexpr, # Padded scale column count (multiple of 8)
SCALE_M_PAD: tl.constexpr, # Padded scale row count (multiple of 256)
):
pid_m = tl.program_id(0)
start_n = tl.program_id(1) * NUM_ITER
# Cast strides to int64 in case M*N > max int32
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)
stride_bs_m = tl.cast(stride_bs_m_in, tl.int64)
stride_bs_n = tl.cast(stride_bs_n_in, tl.int64)
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):
# ── Load input tile ──
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
)
# ── Quantize to MXFP4 (reuses aiter's _mxfp4_quant_op) ──
out_tensor, bs_e8m0 = _mxfp4_quant_op(
x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE
)
# ── Store quantized FP4 data (unchanged from original) ──
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 block scales ──
bs_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
bs_offs_n = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
num_bs_cols = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
if SHUFFLE:
# ── Shuffled store ──
# Copied from aiter's _fused_rms_mxfp4_quant_kernel (lines 173-191).
# Decomposes the 2D scale index [m, n] into sub-indices matching
# gemm_a4w4's (16,16) tile access pattern:
# m → [m//32, m%32//16, m%16] (outer block, half-block, inner row)
# n → [n//8, n%8//4, n%4 ] (outer block, half-block, inner col)
# Then interleaves them so that scales for the same GEMM tile are
# contiguous in memory.
bs_offs_0 = bs_offs_m[:, None] // 32 # which 32-row block
bs_offs_1 = bs_offs_m[:, None] % 32
bs_offs_2 = bs_offs_1 % 16 # row within 16-row half
bs_offs_1 = bs_offs_1 // 16 # which 16-row half (0 or 1)
bs_offs_3 = bs_offs_n[None, :] // 8 # which 8-col block
bs_offs_4 = bs_offs_n[None, :] % 8
bs_offs_5 = bs_offs_4 % 4 # col within 4-col half
bs_offs_4 = bs_offs_4 // 4 # which 4-col half (0 or 1)
bs_offs = (
bs_offs_1
+ bs_offs_4 * 2
+ bs_offs_2 * 2 * 2
+ bs_offs_5 * 2 * 2 * 16
+ bs_offs_3 * 2 * 2 * 16 * 4
+ bs_offs_0 * 2 * 16 * SCALE_N_PAD
)
# Out-of-bounds scales get value 127 (e8m0 for scale=1.0, i.e. no scaling)
bs_valid_mask = (bs_offs_m < M)[:, None] & (bs_offs_n < num_bs_cols)[None, :]
bs_e8m0 = tl.where(bs_valid_mask, bs_e8m0, 127)
# Store mask covers the full padded region
bs_mask = (bs_offs_m < SCALE_M_PAD)[:, None] & (
bs_offs_n < SCALE_N_PAD
)[None, :]
tl.store(bs_ptr + bs_offs, bs_e8m0.to(bs_ptr.type.element_ty), mask=bs_mask)
else:
# ── Linear store (original behavior) ──
bs_offs = (
bs_offs_m[:, None] * stride_bs_m + bs_offs_n[None, :] * stride_bs_n
)
if EVEN_M_N:
tl.store(bs_ptr + bs_offs, bs_e8m0)
else:
bs_mask = (bs_offs_m < M)[:, None] & (
bs_offs_n < num_bs_cols
)[None, :]
tl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask)
def dynamic_mxfp4_quant_shuffled(
x: torch.Tensor,
shuffle: bool = True,
out_fp4: torch.Tensor = None,
out_scale: torch.Tensor = None,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
MXFP4 quantization with optional inline scale shuffle.
Drop-in replacement for:
x_fp4, bs = dynamic_mxfp4_quant(x)
if shuffle:
bs = e8m0_shuffle(bs)
When shuffle=True, the Triton kernel writes scales directly in the
permuted layout that gemm_a4w4 expects. This eliminates:
- The padded tensor allocation in e8m0_shuffle
- The copy into the padded tensor
- The permute + .contiguous() (a full read+write of the scale tensor)
Args:
x: [M, K] bf16/fp16 input tensor.
shuffle: If True, write scales in gemm_a4w4's shuffled layout.
out_fp4: Optional pre-allocated [M, K//2] uint8 output buffer.
out_scale: Optional pre-allocated scale output buffer.
Returns:
(x_fp4, blockscale_e8m0) — same as dynamic_mxfp4_quant, but with
scales already shuffled when shuffle=True.
"""
M, N = x.shape
assert (N // 2) % 2 == 0
MXFP4_QUANT_BLOCK_SIZE = 32
if out_fp4 is not None:
x_fp4 = out_fp4
else:
x_fp4 = torch.empty((M, N // 2), dtype=torch.uint8, device=x.device)
# Scale dimensions
SCALE_N_valid = triton.cdiv(N, MXFP4_QUANT_BLOCK_SIZE)
if shuffle:
# Pad scale dims to multiples required by gemm_a4w4's tile layout:
# rows → multiple of 256 (for 32-row blocks × 8 unroll)
# cols → multiple of 8 (for 8-col blocks)
SCALE_M = triton.cdiv(M, 256) * 256
SCALE_N = triton.cdiv(SCALE_N_valid, 8) * 8
else:
SCALE_M = M
SCALE_N = SCALE_N_valid
if out_scale is not None:
blockscale_e8m0 = out_scale
else:
blockscale_e8m0 = torch.empty(
(SCALE_M, SCALE_N), dtype=torch.uint8, device=x.device
)
# ── Kernel launch config (same heuristics as aiter's dynamic_mxfp4_quant) ──
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
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))
grid = (
triton.cdiv(M, BLOCK_SIZE_M),
triton.cdiv(N, BLOCK_SIZE_N * NUM_ITER),
)
_dynamic_mxfp4_quant_shuffle_kernel[grid](
x,
x_fp4,
blockscale_e8m0,
*x.stride(),
*x_fp4.stride(),
*blockscale_e8m0.stride(),
M=M,
N=N,
MXFP4_QUANT_BLOCK_SIZE=MXFP4_QUANT_BLOCK_SIZE,
SCALING_MODE=0,
NUM_ITER=NUM_ITER,
BLOCK_SIZE_M=BLOCK_SIZE_M,
BLOCK_SIZE_N=BLOCK_SIZE_N,
NUM_STAGES=NUM_STAGES,
# Shuffle params
SHUFFLE=shuffle,
SCALE_N_PAD=SCALE_N,
SCALE_M_PAD=SCALE_M,
# Triton runtime config
num_warps=NUM_WARPS,
waves_per_eu=0,
num_stages=1,
)
return x_fp4, blockscale_e8m0
# ─────────────────────────────────────────────────────────────────────────────
# Kernel implementations
# ─────────────────────────────────────────────────────────────────────────────
def ref_kernel(data: input_t) -> output_t:
"""
Reference: MXFP4 per-1x32 quant on A; B_shuffle, B_scale_sh from generate_input.
gemm_a4w4 with bpreshuffle=True.
"""
import aiter
from aiter import QuantType, dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
def _quant_mxfp4(x, shuffle=True):
x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
if shuffle:
bs_e8m0 = e8m0_shuffle(bs_e8m0)
return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
B = B.contiguous()
m, k = A.shape
n, _ = B.shape
A_q, A_scale_sh = _quant_mxfp4(A, shuffle=True)
out_gemm = aiter.gemm_a4w4(
A_q,
B_shuffle,
A_scale_sh,
B_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
return out_gemm
def custom_kernel1(data: input_t) -> output_t:
"""
Optimized MXFP4 quant + GEMM using aiter's fused_rms_mxfp4_quant.
Key optimization vs custom_ref_kernel:
─────────────────────────────────────
Reference does 3 steps (3 kernel launches, 2 unnecessary global memory round-trips):
1. dynamic_mxfp4_quant(A) → Triton kernel: writes A_q + A_scale to global mem
2. e8m0_shuffle(A_scale) → PyTorch ops (alloc padded tensor, copy, view,
permute, .contiguous()) = extra kernel + full
read+write of scale tensor through global memory
3. gemm_a4w4(...) → CK/ASM GEMM: reads A_q + A_scale back
This version does 2 steps (2 kernel launches):
1. fused_rms_mxfp4_quant(A, shuffle=True)
→ Single Triton kernel that performs:
a) RMSNorm (identity when weight=ones, eps=0,
and input is already unit-RMS — see note below)
b) MXFP4 quantization
c) e8m0 scale shuffle (inline, no extra alloc)
All in one pass over A, writing shuffled scales
directly without the pad+permute+contiguous dance.
2. gemm_a4w4(...) → CK/ASM GEMM (unchanged)
IMPORTANT NOTE on correctness:
─────────────────────────────
fused_rms_mxfp4_quant always applies RMSNorm: x_out = x / rms(x) * weight.
With weight=ones and eps=0, this normalizes each row to unit RMS norm.
This CHANGES the magnitude of A, so the GEMM output will differ from the
reference by a per-row scaling factor (rms of each row of A).
For a truly correct drop-in replacement, we would need either:
a) A patched dynamic_mxfp4_quant that accepts shuffle=True (the underlying
Triton kernel _dynamic_mxfp4_quant_kernel does NOT have a SHUFFLE param), or
b) A standalone Triton wrapper that calls _mxfp4_quant_op + inline shuffle
without the RMSNorm.
This implementation demonstrates the fused kernel pattern. If correctness vs
the reference is required, fall back to custom_ref_kernel until (a) or (b)
is available.
"""
import aiter
from aiter import dtypes
from aiter.ops.triton.quant.fused_mxfp4_quant import fused_rms_mxfp4_quant
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
m, k = A.shape
# --- Step 1: Fused quant + shuffle in a single Triton kernel ---
# fused_rms_mxfp4_quant does: RMSNorm → MXFP4 quant → scale shuffle
# We pass weight=ones and eps=0 so RMSNorm becomes x / rms(x).
# The shuffle=True flag writes e8m0 scales in the permuted layout that
# gemm_a4w4 expects, avoiding the separate e8m0_shuffle() call which
# would allocate a padded tensor, copy, view(6D), permute, .contiguous().
ones = torch.ones(k, dtype=A.dtype, device=A.device)
(A_q, A_scale_sh), _, _, _ = fused_rms_mxfp4_quant(
A,
x1_weight=ones,
x1_epsilon=0.0,
shuffle=True, # <-- inline scale shuffle, no separate kernel
)
# Reinterpret raw uint8 outputs as the fp4x2 / fp8_e8m0 dtypes
# that gemm_a4w4 expects (these are zero-cost view operations, no copy).
A_q = A_q.view(dtypes.fp4x2)
A_scale_sh = A_scale_sh.view(dtypes.fp8_e8m0)
# --- Step 2: GEMM (unchanged from reference) ---
# gemm_a4w4 dispatches to CK (Composable Kernel) or hand-tuned ASM
# depending on shape and tuning config. bpreshuffle=True tells it that
# B is already in (16,16)-tile-coalesced layout from shuffle_weight().
out_gemm = aiter.gemm_a4w4(
A_q,
B_shuffle,
A_scale_sh,
B_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
return out_gemm
def custom_kernel2(data: input_t) -> output_t:
"""
Optimized MXFP4 quant + GEMM — correct drop-in for ref_kernel.
Uses a patched dynamic_mxfp4_quant that bakes the e8m0 scale shuffle
into the Triton kernel's store logic. This eliminates the separate
e8m0_shuffle() call (padded alloc + copy + permute + .contiguous()).
Reference (3 kernel launches):
1. dynamic_mxfp4_quant(A) → Triton: quant, writes A_q + A_scale
2. e8m0_shuffle(A_scale) → PyTorch: alloc padded buf, copy, permute, .contiguous()
3. gemm_a4w4(...) → CK/ASM GEMM
This version (2 kernel launches):
1. dynamic_mxfp4_quant_shuffled(A, shuffle=True)
→ Triton: quant + shuffled scale store in one kernel
2. gemm_a4w4(...) → CK/ASM GEMM (unchanged)
"""
import aiter
from aiter import dtypes
A, B, B_q, B_shuffle, B_scale_sh = data
m, k = A.shape
# Step 1: Quantize A with inline scale shuffle (single Triton kernel).
# Reuses aiter's _mxfp4_quant_op for the FP4 math, but writes
# e8m0 block scales directly in gemm_a4w4's permuted layout.
# No separate e8m0_shuffle needed.
A_q, A_scale_sh = dynamic_mxfp4_quant_shuffled(A, shuffle=True)
# Step 2: GEMM (unchanged).
# .view(dtypes.fp4x2) / .view(dtypes.fp8_e8m0) are zero-cost reinterprets.
out = aiter.gemm_a4w4(
A_q.view(dtypes.fp4x2),
B_shuffle,
A_scale_sh.view(dtypes.fp8_e8m0),
B_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
return out
def custom_kernel_single_fused(data: input_t) -> output_t:
"""
Single-kernel fused quant+GEMM using gemm_a16wfp4_preshuffle.
This eliminates the separate A quantization kernel entirely.
The preshuffle GEMM kernel quantizes bf16 A to MXFP4 on-the-fly
in registers during the GEMM loop (via _mxfp4_quant_op + tl.dot_scaled),
so there are no intermediate FP4 buffers for A written to global memory.
Reference (3 kernel launches):
1. dynamic_mxfp4_quant(A) → writes A_q + A_scale to global mem
2. e8m0_shuffle(A_scale) → alloc padded buf, permute, contiguous
3. gemm_a4w4(...) → CK/ASM GEMM
custom_kernel2 (2 kernel launches):
1. dynamic_mxfp4_quant_shuffled(A) → quant + shuffled scale store
2. gemm_a4w4(...) → CK/ASM GEMM
This version (1 kernel launch + optional splitK reduce):
1. gemm_a16wfp4_preshuffle(A, B, B_scales)
→ Triton GEMM that quantizes A per-tile in registers
→ For small M, uses splitK (parallel K reduction) for better CU utilization
B format notes:
- B_shuffle [N, K//2] from shuffle_weight(B_q, (16,16)) must be reshaped
to [N//16, K//2*16] — same data, different view matching the kernel's
tiled B loading pattern.
- B_scale_sh [padded_N, padded_K_scale] from e8m0_shuffle must be reshaped
to [padded_N//32, padded_K_scale*32] — same shuffled data, different 2D
view matching the kernel's scale pointer arithmetic. The kernel un-shuffles
scales in registers (reshape+permute = inverse of e8m0_shuffle).
"""
A, B, B_q, B_shuffle, B_scale_sh = data
m, k = A.shape
n = B.shape[0]
# Reshape B: [N, K//2] → [N//16, K//2*16]
B_w = B_shuffle.view(torch.uint8).reshape(n // 16, (k // 2) * 16)
# Reshape B scales: [padded_N, padded_K_scale] → [padded_N//32, padded_K_scale*32]
bs = B_scale_sh.view(torch.uint8)
B_scale_w = bs.reshape(bs.shape[0] // 32, bs.shape[1] * 32)
configs = {
# K=512: NUM_KSPLIT=1 (BSK covers full K), no atomic needed
(2880, 512): {
"BLOCK_SIZE_M": 8 if m <= 8 else 32,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 4 if m <= 8 else 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
(4096, 512): {
"BLOCK_SIZE_M": 16 if m <= 32 else 32,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 4 if m <= 32 else 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
# K=7168: splitK=14, atomic eliminates reduce kernel
(2112, 7168): {
"BLOCK_SIZE_M": 8 if m <= 8 else (16 if m <= 64 else 32),
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 1 if m <= 8 else 4,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 14,
},
# K=2048: splitK=4, atomic eliminates reduce kernel
(7168, 2048): {
"BLOCK_SIZE_M": 16 if m <= 64 else 32,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 8 if m >= 128 else 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 4 if m <= 64 else 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 4,
},
# K=1536: splitK=3
(3072, 1536): {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 8 if m >= 128 else 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 3,
},
}
config = configs.get((n, k), None)
# Note: Atomic splitK was attempted (tl.atomic_add to eliminate reduce kernel)
# but was slower due to torch.zeros init cost, atomic contention, and fp32→bf16
# conversion overhead. Standard splitK + reduce kernel remains faster.
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
y = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
y = gemm_a16wfp4_preshuffle(
A, B_w, B_scale_w,
dtype=torch.bfloat16,
y=y,
config=config,
)
return y
def custom_kernel(data: input_t) -> output_t:
"""
Hybrid dispatch: picks the fastest kernel per (m, n, k) shape.
Three paths:
- CK/ASM path (gemm_a4w4): hand-tuned assembly, best for some shapes
- Single-launch Triton (gemm_a16wfp4_preshuffle): fused A quant in registers
- 2-launch Triton (quant + gemm_afp4wfp4_preshuffle): separate quant, FP4×FP4 GEMM
The dispatch table below maps each benchmark shape to its fastest path.
"""
A, B, B_q, B_shuffle, B_scale_sh = data
m, k = A.shape
n = B.shape[0]
# Dispatch key: (m, n, k) for exact match, fallback to (n, k) heuristic
# Path: "single" = gemm_a16wfp4_preshuffle, "dual" = quant + gemm_afp4wfp4_preshuffle, "ck" = quant + gemm_a4w4
dispatch = {
# M=4, N=2880, K=512: single-launch (tiny M, fused quant wins)
(4, 2880, 512): "single",
# M=16, N=2112, K=7168: single-launch (splitK=14, avoids quant launch)
(16, 2112, 7168): "single",
# M=32, N=4096, K=512: single-launch (small K, 1 tile in K)
(32, 4096, 512): "single",
# M=32, N=2880, K=512: single-launch (small K, 1 tile in K)
(32, 2880, 512): "single",
# M=64, N=7168, K=2048: 2-launch Triton (larger M, separate quant is cheaper)
(64, 7168, 2048): "ck",
# M=256, N=3072, K=1536: 2-launch Triton (large M, quant kernel efficient)
(256, 3072, 1536): "ck",
}
path = dispatch.get((m, n, k), None)
if path is None:
# Heuristic fallback: large M → dual, small M → single
path = "dual" if m >= 64 else "single"
if path == "ck":
# ── CK/ASM path: quant A + gemm_a4w4 (hand-tuned assembly) ──
import aiter
from aiter import dtypes
A_q, A_scale_sh = dynamic_mxfp4_quant_shuffled(A, shuffle=True)
out = aiter.gemm_a4w4(
A_q.view(dtypes.fp4x2),
B_shuffle,
A_scale_sh.view(dtypes.fp8_e8m0),
B_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
return out
# Both Triton paths need reshaped B
B_w = B_shuffle.view(torch.uint8).reshape(n // 16, (k // 2) * 16)
bs = B_scale_sh.view(torch.uint8)
B_scale_w = bs.reshape(bs.shape[0] // 32, bs.shape[1] * 32)
if path == "dual":
# ── 2-launch path: quant A separately, then FP4×FP4 GEMM ──
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import gemm_afp4wfp4_preshuffle
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import _get_config
configs_2launch = {
(7168, 2048): {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 1,
"num_warps": 2,
"num_stages": 2,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 2,
},
(3072, 1536): {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 3,
},
}
config = configs_2launch.get((n, k), None)
if config is None:
k_internal = k // 2
config, _ = _get_config(m, n, k_internal, shuffle=True)
block_m = config["BLOCK_SIZE_M"]
if block_m >= 32 and m >= 32:
A_q, A_scale = dynamic_mxfp4_quant_shuffled(A, shuffle=True)
A_scale_w = A_scale.reshape(A_scale.shape[0] // 32, A_scale.shape[1] * 32)
else:
A_q, A_scale = dynamic_mxfp4_quant_shuffled(A, shuffle=False)
A_scale_w = A_scale
y = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
y = gemm_afp4wfp4_preshuffle(
A_q, B_w, A_scale_w, B_scale_w,
dtype=torch.bfloat16,
y=y,
config=config,
)
return y
else:
# ── Single-launch path: fused quant+GEMM ──
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
configs_single = {
(2880, 512): {
"BLOCK_SIZE_M": 8 if m <= 8 else 32,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 4 if m <= 8 else 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
(4096, 512): {
"BLOCK_SIZE_M": 16 if m <= 32 else 32,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 4 if m <= 32 else 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
(2112, 7168): {
"BLOCK_SIZE_M": 8 if m <= 8 else (16 if m <= 64 else 32),
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 1 if m <= 8 else 4,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 14,
},
(7168, 2048): {
"BLOCK_SIZE_M": 16 if m <= 64 else 32,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 8 if m >= 128 else 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 4 if m <= 64 else 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 4,
},
(3072, 1536): {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 8 if m >= 128 else 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 3,
},
}
config = configs_single.get((n, k), None)
y = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
y = gemm_a16wfp4_preshuffle(
A, B_w, B_scale_w,
dtype=torch.bfloat16,
y=y,
config=config,
)
return y
def custom_kernel3(data: input_t) -> output_t:
"""
Optimized 2-launch GEMM using gemm_afp4wfp4_preshuffle (Triton FP4×FP4).
Optimizations vs previous version:
1. For M >= 32: use shuffle=True in quant kernel to produce preshuffled
A scales directly, avoiding a separate .permute().contiguous() launch
(saves 1 kernel launch = ~3-5us)
2. Custom per-shape configs with tuned NUM_KSPLIT for better CU utilization
3. Pre-allocated output tensor passed to GEMM to avoid torch.empty overhead
"""
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import gemm_afp4wfp4_preshuffle
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import _get_config
A, B, B_q, B_shuffle, B_scale_sh = data
m, k = A.shape
n = B.shape[0]
# Reshape B for preshuffle kernel: [N, K//2] → [N//16, K//2*16]
B_w = B_shuffle.view(torch.uint8).reshape(n // 16, (k // 2) * 16)
# Reshape B scales: [padded_N, padded_K_scale] → [padded_N//32, padded_K_scale*32]
bs = B_scale_sh.view(torch.uint8)
B_scale_w = bs.reshape(bs.shape[0] // 32, bs.shape[1] * 32)
# Per-shape configs. Shapes with tuned JSON configs (N=2112,K=7168;
# N=4096,K=512; N=3072,K=1536) use None → _get_config loads the JSON.
# Shapes without tuned configs get custom configs here.
configs = {
# N=2880, K=512: K_int=256, can't splitK. Use BSN=32 for more tiles.
# M=4: grid=1*90=90, M=32: grid=1*90=90 (BSM=32,BSN=32)
(2880, 512): {
"BLOCK_SIZE_M": 8 if m < 32 else 32,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 1,
"num_warps": 2,
"num_stages": 2,
"waves_per_eu": 4 if m <= 8 else 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
},
# N=7168, K=2048: K_int=1024. splitK=2 doubles grid tiles.
# M=64: grid=2*2*112=448 tiles (vs 224 without splitK)
(7168, 2048): {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 1,
"num_warps": 2,
"num_stages": 2,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 2,
},
}
config = configs.get((n, k), None)
if config is None:
# Use tuned JSON config (exists for N=2112,K=7168; N=4096,K=512; N=3072,K=1536)
k_internal = k // 2
config, _ = _get_config(m, n, k_internal, shuffle=True)
block_m = config["BLOCK_SIZE_M"]
# Quantize A with scale format matching what the GEMM kernel expects.
# Key optimization: for block_m >= 32, we use shuffle=True to produce
# preshuffled scales directly in the quant kernel, avoiding a separate
# .permute().contiguous() kernel launch.
if block_m >= 32 and m >= 32:
# Shuffle=True: scales written in e8m0_shuffle layout [pad_M, pad_N]
A_q, A_scale = dynamic_mxfp4_quant_shuffled(A, shuffle=True)
# Reshape to [pad_M//32, pad_N*32] — zero-cost view, same data layout
A_scale_w = A_scale.reshape(A_scale.shape[0] // 32, A_scale.shape[1] * 32)
else:
# Linear scales for small M (block_m < 32)
A_q, A_scale = dynamic_mxfp4_quant_shuffled(A, shuffle=False)
A_scale_w = A_scale
# Pre-allocate output
y = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
out = gemm_afp4wfp4_preshuffle(
A_q, B_w, A_scale_w, B_scale_w,
dtype=torch.bfloat16,
y=y,
config=config,
)
return out
scrolls · 851 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