submission 732442
josusanmartin · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 919 lines, June 9 Researcher Reciprocity License v1.0.
submission_v1156.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-732442?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:f35b7daf4334f35b89dcf1a4f1f03a4e34dec1c22f2034c110cbd9fea36f756a
license declaredunknown
license concludedunknown
authorsjosusanmartin
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
"""Convert 2 FP32 values to packed FP4 byte using hardware instruction.num-warps = 4
num_warps=4, num_stages=1,split-k
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitkstages = 1
num_warps=4, num_stages=1,tile-n = 64
ACTUAL_KSPLIT=ACTUAL_KSPLIT, MAX_KSPLIT=MAX_KSPLIT, BLOCK_N=64,Kernel source
submission_v1156.py919 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
v1122_cand0873: exact M=16 XCD + M=16 waves=0 + M=256 stages=2 + M=32 stages=3.
Key changes from cand_0902 base:
- exact M=16 XCD kernel (hardcoded KSPLIT=7, NUM_K_ITER=2, pid decomposition)
- M=16 waves_per_eu=0 (auto, was 2)
- M=256 nosplit stages=2 (was 1)
- M=32 shapes stages=3 (was 4)
- duplicate exact M=16 kernel removed (faster JIT)
Measured benchmark result:
- ~7.40 us geomean (non-ranked)
- [5.98, 8.55, 6.16, 6.00, 9.74, 9.81] us
"""
from __future__ import annotations
import os
os.environ["TRITON_HIP_USE_BLOCK_PINGPONG"] = "0"
import torch
import triton
import triton.language as tl
import aiter
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op as _mxfp4_quant_op_sw
from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid, remap_xcd
from task import input_t, output_t
# ---- Hardware FP4 conversion (prescale + scale=1.0 via VGPR) ----
@triton.jit
def _hw_fp4_convert_pair(val0, val1, scale):
"""Convert 2 FP32 values to packed FP4 byte using hardware instruction.
CK-style: no v_mov_b32, "=v" output constraint, VGPR scale.
"""
return tl.inline_asm_elementwise(
"v_cvt_scalef32_pk_fp4_f32 $0, $1, $2, $3",
"=v,v,v,v",
[val0, val1, scale],
dtype=tl.int32,
is_pure=True,
pack=1,
)
@triton.jit
def _mxfp4_quant_op_hw(
x,
BLOCK_SIZE_N,
BLOCK_SIZE_M,
MXFP4_QUANT_BLOCK_SIZE,
):
"""Hybrid quant: bitwise exponent extraction (no log2/floor) + exp2 prescale.
Saves 2 SFU ops vs original hw quant while keeping low register pressure."""
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
x = x.to(tl.float32).reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
# Step 1: Bitwise exponent extraction (avoids log2 + floor)
amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
amax_bits = amax.to(tl.int32, bitcast=True)
amax_bits = (amax_bits + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
exponent_biased = ((amax_bits >> 23) & 0xFF).to(tl.int32)
# scale_e8m0_unbiased = exponent_biased - 127 - 2 = exponent_biased - 129
scale_e8m0_unbiased = exponent_biased - 129
# tl.clamp doesn't support int32, use tl.where instead
scale_e8m0_unbiased = tl.where(scale_e8m0_unbiased < -127, -127,
tl.where(scale_e8m0_unbiased > 127, 127, scale_e8m0_unbiased))
bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
# Step 2: Prescale via exp2 (keeps low register pressure, avoids bitwise prescale)
prescale = tl.exp2((-scale_e8m0_unbiased).to(tl.float32))
x = x * tl.broadcast_to(prescale, x.shape)
# Step 3: Convert prescaled values to FP4 using hw instruction with scale=1.0
x_pairs = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2)
evens, odds = tl.split(x_pairs)
ones = tl.full(evens.shape, 1.0, dtype=tl.float32)
packed_i32 = _hw_fp4_convert_pair(evens, odds, ones)
x_fp4 = (packed_i32 & 0xFF).to(tl.uint8)
x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
@triton.jit
def _mxfp4_quant_op_hw_bitwise(
x,
BLOCK_SIZE_N,
BLOCK_SIZE_M,
MXFP4_QUANT_BLOCK_SIZE,
):
"""Full bitwise quant — avoids log2, floor, AND exp2 via IEEE 754 manipulation."""
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
x = x.to(tl.float32).reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
# Step 1: Compute per-block scale via bitwise exponent extraction
amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
amax_bits = amax.to(tl.int32, bitcast=True)
amax_bits = (amax_bits + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
# Extract IEEE 754 biased exponent
exponent_biased = ((amax_bits >> 23) & 0xFF).to(tl.int32)
# bs_e8m0 = exponent_biased - 2, clamped to [0, 254]
bs_e8m0_i32 = tl.where(exponent_biased < 2, 0, tl.where(exponent_biased > 256, 254, exponent_biased - 2))
bs_e8m0 = bs_e8m0_i32.to(tl.uint8)
# Step 2: Construct prescale as 2^(-scale_e8m0_unbiased) via bit manipulation
# prescale_exponent_biased = 127 - scale_e8m0_unbiased = 127 - (exponent_biased - 129) = 256 - exponent_biased
# After clamping: prescale_exponent_biased = 127 - (bs_e8m0_i32 - 127) = 254 - bs_e8m0_i32
prescale_exp = 254 - bs_e8m0_i32
prescale_bits = (prescale_exp << 23)
prescale = prescale_bits.to(tl.float32, bitcast=True)
x = x * tl.broadcast_to(prescale, x.shape)
# Step 3: HW convert with scale=1.0
x_pairs = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2)
evens, odds = tl.split(x_pairs)
ones = tl.full(evens.shape, 1.0, dtype=tl.float32)
packed_i32 = _hw_fp4_convert_pair(evens, odds, ones)
x_fp4 = (packed_i32 & 0xFF).to(tl.uint8)
x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
# Monkey-patch BITWISE quant for preshuffle kernel path (faster for M=64, M=256)
import aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 as _kernel_module
_kernel_module._mxfp4_quant_op = _mxfp4_quant_op_hw_bitwise
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import _gemm_afp4wfp4_reduce_kernel
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
# ---- Enhanced preshuffle kernel: +XCD remap, +.wt store, +accumulator arg ----
@triton.heuristics({
"EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)
and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),
"GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])
* triton.cdiv(args["N"], args["BLOCK_SIZE_N"]),
})
@triton.jit
def _enhanced_preshuffle_kernel(
a_ptr, b_ptr, c_ptr, b_scales_ptr,
M, N, K,
stride_am, stride_ak, stride_bn, stride_bk,
stride_ck, stride_cm, stride_cn, stride_bsn, stride_bsk,
BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr,
NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,
EVEN_K: tl.constexpr,
num_warps: tl.constexpr, num_stages: tl.constexpr,
waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
GRID_MN: tl.constexpr, PREQUANT: tl.constexpr, cache_modifier: tl.constexpr,
):
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_bsk > 0)
tl.assume(stride_bsn > 0)
pid_unified = tl.program_id(axis=0)
pid_k = pid_unified % NUM_KSPLIT
pid = pid_unified // NUM_KSPLIT
# XCD remap ONLY for small grids ≤90 blocks (K=512 cases)
# Hurts M=64 (224 blocks) and M=256 (384 blocks)
# no remap
num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
if NUM_KSPLIT == 1:
pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
else:
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
SCALE_GROUP_SIZE: tl.constexpr = 32
if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16
offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak)
offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
b_ptrs = b_ptr + (offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk)
offs_bsn = (pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, (BLOCK_SIZE_N // 32))) % N
offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)
b_scale_ptrs = b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
b_scales = (tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
.reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8, 4, 16, 2, 2, 1)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE))
if EVEN_K:
a_bf16 = tl.load(a_ptrs, cache_modifier=".ca")
b = tl.load(b_ptrs, cache_modifier=cache_modifier)
b = (b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
.trans(1, 0))
if PREQUANT:
a, a_scales = _mxfp4_quant_op_hw_bitwise(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
# CHANGE 2: Pass accumulator as arg (enables HW FMA fusion)
accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)
a_ptrs += BLOCK_SIZE_K * stride_ak
b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
b_scale_ptrs += BLOCK_SIZE_K * stride_bsk
c = accumulator.to(c_ptr.type.element_ty)
offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :] + pid_k * stride_ck
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
if GRID_MN <= 224:
tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")
else:
tl.store(c_ptrs, c, mask=c_mask)
@triton.heuristics({
"EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0),
"GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])
})
@triton.jit
def _enhanced_preshuffle_exact_m64_kernel(
a_ptr, b_ptr, c_ptr, b_scales_ptr,
M, N, K,
stride_am, stride_ak, stride_bn, stride_bk,
stride_ck, stride_cm, stride_cn, stride_bsn, stride_bsk,
BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr,
NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,
EVEN_K: tl.constexpr,
num_warps: tl.constexpr, num_stages: tl.constexpr,
waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
GRID_MN: tl.constexpr, PREQUANT: tl.constexpr, cache_modifier: tl.constexpr,
):
pid = tl.program_id(axis=0)
pid_m = pid // 56
pid_n = pid % 56
SCALE_GROUP_SIZE: tl.constexpr = 32
NUM_K_ITER: tl.constexpr = 4
offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_k_bf16[None, :] * stride_ak
offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
offs_bn = pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)
b_ptrs = b_ptr + offs_bn[:, None] * stride_bn + offs_k_shuffle_arr[None, :] * stride_bk
offs_bsn = pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)
offs_ks = tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)
b_scale_ptrs = b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for _ in range(NUM_K_ITER):
b_scales = (
tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
.reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8, 4, 16, 2, 2, 1)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
)
a_bf16 = tl.load(a_ptrs, cache_modifier=".ca")
b = tl.load(b_ptrs, cache_modifier=cache_modifier)
b = (
b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
.trans(1, 0)
)
if PREQUANT:
a, a_scales = _mxfp4_quant_op_hw_bitwise(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)
a_ptrs += BLOCK_SIZE_K * stride_ak
b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
b_scale_ptrs += BLOCK_SIZE_K * stride_bsk
c = accumulator.to(c_ptr.type.element_ty)
offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :]
tl.store(c_ptrs, c, cache_modifier=".wt")
@triton.heuristics({
"EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0),
"GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])
})
@triton.jit
def _enhanced_preshuffle_exact_m256_kernel(
a_ptr, b_ptr, c_ptr, b_scales_ptr,
M, N, K,
stride_am, stride_ak, stride_bn, stride_bk,
stride_ck, stride_cm, stride_cn, stride_bsn, stride_bsk,
BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr,
NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,
EVEN_K: tl.constexpr,
num_warps: tl.constexpr, num_stages: tl.constexpr,
waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
GRID_MN: tl.constexpr, PREQUANT: tl.constexpr, cache_modifier: tl.constexpr,
):
pid = tl.program_id(axis=0)
pid_m = pid // 24
pid_n = pid % 24
SCALE_GROUP_SIZE: tl.constexpr = 32
NUM_K_ITER: tl.constexpr = 6
offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_k_bf16[None, :] * stride_ak
offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
offs_bn = pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)
b_ptrs = b_ptr + offs_bn[:, None] * stride_bn + offs_k_shuffle_arr[None, :] * stride_bk
offs_bsn = pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)
offs_ks = tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)
b_scale_ptrs = b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for _ in range(NUM_K_ITER):
b_scales = (
tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
.reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8, 4, 16, 2, 2, 1)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
)
a_bf16 = tl.load(a_ptrs, cache_modifier=".ca")
b = tl.load(b_ptrs, cache_modifier=cache_modifier)
b = (
b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
.trans(1, 0)
)
if PREQUANT:
a, a_scales = _mxfp4_quant_op_hw_bitwise(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)
a_ptrs += BLOCK_SIZE_K * stride_ak
b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
b_scale_ptrs += BLOCK_SIZE_K * stride_bsk
c = accumulator.to(c_ptr.type.element_ty)
offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :]
tl.store(c_ptrs, c)
# ---- Exact M=16 XCD kernel (hardcoded for M=16, N=2112, KSPLIT=7) ----
@triton.heuristics({"EVEN_K": lambda args: True, "GRID_MN": lambda args: 34})
@triton.jit
def _gemm_exact_m16_xcd_kernel(
a_ptr, b_ptr, c_ptr, b_scales_ptr,
M, N, K,
stride_am, stride_ak, stride_bn, stride_bk,
stride_ck, stride_cm, stride_cn,
stride_bsn, stride_bsk,
BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr,
NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,
EVEN_K: tl.constexpr,
num_warps: tl.constexpr, num_stages: tl.constexpr,
waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
GRID_MN: tl.constexpr,
PREQUANT: tl.constexpr, cache_modifier: tl.constexpr,
):
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_bsk > 0)
tl.assume(stride_bsn > 0)
pid_linear = tl.program_id(axis=0)
pid_k = pid_linear % 7
pid = remap_xcd(pid_linear // 7, 34, NUM_XCDS=8)
pid_m = pid % 2
pid_n = pid // 2
SCALE_GROUP_SIZE: tl.constexpr = 32
NUM_K_ITER: tl.constexpr = 2
offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16
offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak
offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
offs_bn = pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)
b_ptrs = b_ptr + offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk
offs_bsn = pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)
offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)
b_scale_ptrs = b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for _ in range(NUM_K_ITER):
b_scales = (
tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
.reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8, 4, 16, 2, 2, 1)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
)
a_bf16 = tl.load(a_ptrs, cache_modifier=".ca")
b = tl.load(b_ptrs, cache_modifier=cache_modifier)
b = (
b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
.trans(1, 0)
)
a, a_scales = _mxfp4_quant_op_hw(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)
a_ptrs += BLOCK_SIZE_K * stride_ak
b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
b_scale_ptrs += BLOCK_SIZE_K * stride_bsk
c = accumulator.to(c_ptr.type.element_ty)
offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :] + pid_k * stride_ck
c_mask = offs_cn[None, :] < N
tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")
_BF16 = dtypes.bf16
_FP4X2 = dtypes.fp4x2
_FP8_E8M0 = dtypes.fp8_e8m0
_KERNEL_32X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_FUSED_CONFIGS = {
(4, 2880, 512): {
"BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3,
"waves_per_eu": 2, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": 1,
},
(16, 2112, 7168): {
"BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 2, "num_warps": 4, "num_stages": 2,
"waves_per_eu": 0, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": 7,
},
(32, 4096, 512): {
"BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3,
"waves_per_eu": 1, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": 1,
},
(32, 2880, 512): {
"BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3,
"waves_per_eu": 1, "matrix_instr_nonkdim": 16,
"cache_modifier": None, "NUM_KSPLIT": 1,
},
# M=256 via fused preshuffle — v925: BM=16 + cache=".cg"
(256, 3072, 1536): {
"BLOCK_SIZE_M": 16, "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": 1,
},
}
_FUSED_K7168_KSPLIT14 = {
"BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3,
"waves_per_eu": 1, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": 14,
}
_FUSED_K1536_SPLIT = {
"BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3,
"waves_per_eu": 1, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": 3,
}
_FUSED_K1536_NOSPLIT = {
"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
"waves_per_eu": 1, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": 1,
}
_FUSED_BM32_K512 = {
"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1, "num_warps": 8, "num_stages": 3,
"waves_per_eu": 1, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": 1,
}
_FUSED_K2048_NOSPLIT = {
"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
"waves_per_eu": 1, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": 1,
}
_ASM_CONFIGS = {
# Removed M=256 from ASM — use fused preshuffle instead
}
_QUANT_BLOCK = 32
_QUANT_TILE = 128
# ---- Standalone quant kernel for ASM path (SOFTWARE quant) ----
@triton.jit
def _quant_kernel_asm_layout(
x_ptr, x_fp4_ptr, bs_ptr,
stride_x_m, stride_x_n,
stride_x_fp4_m, stride_x_fp4_n,
stride_bs_m, stride_bs_n,
M: tl.constexpr, N: tl.constexpr,
scaleN: tl.constexpr, scaleM_pad: tl.constexpr, scaleN_pad: tl.constexpr,
BLOCK_SIZE: tl.constexpr, MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
SHUFFLE: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
stride_x_m = tl.cast(stride_x_m, tl.int64)
stride_x_n = tl.cast(stride_x_n, tl.int64)
stride_x_fp4_m = tl.cast(stride_x_fp4_m, tl.int64)
stride_x_fp4_n = tl.cast(stride_x_fp4_n, tl.int64)
x_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
x_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE)
x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
x = tl.load(x_ptr + x_offs, mask=x_mask).to(tl.float32)
# Use SOFTWARE quant for ASM path
out_tensor, bs_e8m0 = _mxfp4_quant_op_sw(
x, MXFP4_QUANT_BLOCK_SIZE, BLOCK_SIZE, MXFP4_QUANT_BLOCK_SIZE,
)
out_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
out_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE // 2 + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE // 2)
out_offs = out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
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)
bs_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
bs_offs_n = pid_n
if SHUFFLE:
bs_offs_0 = bs_offs_m[:, None] // 32
bs_offs_1 = bs_offs_m[:, None] % 32
bs_offs_2 = bs_offs_1 % 16
bs_offs_1 = bs_offs_1 // 16
bs_offs_3 = bs_offs_n[None, :] // 8
bs_offs_4 = bs_offs_n[None, :] % 8
bs_offs_5 = bs_offs_4 % 4
bs_offs_4 = bs_offs_4 // 4
bs_offs = (
bs_offs_1 + bs_offs_4 * 2 + bs_offs_2 * 4
+ bs_offs_5 * 64 + bs_offs_3 * 256
+ bs_offs_0 * 32 * scaleN
)
bs_mask1 = (bs_offs_m < M)[:, None] & (bs_offs_n < scaleN)[None, :]
bs_mask2 = (bs_offs_m < scaleM_pad)[:, None] & (bs_offs_n < scaleN_pad)[None, :]
bs_e8m0 = tl.where(bs_mask1, bs_e8m0, 127)
tl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask2)
else:
bs_offs = bs_offs_m[:, None] * stride_bs_m + bs_offs_n[None, :] * stride_bs_n
bs_mask = (bs_offs_m < M)[:, None] & (bs_offs_n < N)[None, :]
tl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask)
# ---- Optimized XCD kernel (fused, uses hw quant) ----
@triton.heuristics(
{
"EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)
and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),
"GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])
* triton.cdiv(args["N"], args["BLOCK_SIZE_N"]),
}
)
@triton.jit
def _gemm_optimized_xcd_kernel(
a_ptr, b_ptr, c_ptr, b_scales_ptr,
M, N, K,
stride_am, stride_ak, stride_bn, stride_bk,
stride_ck, stride_cm, stride_cn,
stride_bsn, stride_bsk,
BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr,
NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,
EVEN_K: tl.constexpr,
num_warps: tl.constexpr, num_stages: tl.constexpr,
waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
GRID_MN: tl.constexpr,
PREQUANT: tl.constexpr, cache_modifier: tl.constexpr,
):
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_bsk > 0)
tl.assume(stride_bsn > 0)
pid_linear = tl.program_id(axis=0)
pid_k = pid_linear % NUM_KSPLIT
pid = remap_xcd(pid_linear // NUM_KSPLIT, GRID_MN, NUM_XCDS=8)
num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
SCALE_GROUP_SIZE: tl.constexpr = 32
if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16
offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
a_ptrs = a_ptr + (
offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak
)
offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
b_ptrs = b_ptr + (
offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk
)
offs_bsn = (
pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, (BLOCK_SIZE_N // 32))
) % N
offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(
0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32
)
b_scale_ptrs = (
b_scales_ptr
+ offs_bsn[:, None] * stride_bsn
+ offs_ks[None, :] * stride_bsk
)
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for _ in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
b_scales = (
tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
.reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8, 4, 16, 2, 2, 1)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
)
if EVEN_K:
a_bf16 = tl.load(a_ptrs, cache_modifier=".ca")
b = tl.load(b_ptrs, cache_modifier=cache_modifier)
b = (
b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
.trans(1, 0)
)
if PREQUANT:
a, a_scales = _mxfp4_quant_op_hw(
a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32
)
accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)
a_ptrs += BLOCK_SIZE_K * stride_ak
b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
b_scale_ptrs += BLOCK_SIZE_K * stride_bsk
c = accumulator.to(c_ptr.type.element_ty)
offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
c_ptrs = (
c_ptr
+ stride_cm * offs_cm[:, None]
+ stride_cn * offs_cn[None, :]
+ pid_k * stride_ck
)
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
if GRID_MN <= 224:
tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")
else:
tl.store(c_ptrs, c, mask=c_mask)
def _e8m0_shuffle_safe(scale: torch.Tensor) -> torch.Tensor:
m, n = scale.shape
scale_padded = torch.empty(
((m + 255) // 256) * 256, ((n + 7) // 8) * 8,
dtype=scale.dtype, device=scale.device,
)
scale_padded.fill_(0x7F)
scale_padded[:m, :n] = scale
sm, sn = scale_padded.shape
return (
scale_padded.view(sm // 32, 2, 16, sn // 8, 2, 4)
.permute(0, 3, 5, 2, 4, 1).contiguous().view(sm, sn)
)
@triton.jit
def _reduce_m16_splitk7_kernel(
y_ptr, out_ptr,
N,
stride_yk, stride_ym, stride_yn,
stride_om, stride_on,
ACTUAL_KSPLIT: tl.constexpr,
MAX_KSPLIT: tl.constexpr,
BLOCK_N: tl.constexpr,
):
pid_n = tl.program_id(axis=0)
offs_m = tl.arange(0, 16)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
n_mask = offs_n < N
acc = tl.zeros((16, BLOCK_N), dtype=tl.float32)
for ks in range(MAX_KSPLIT):
if ks < ACTUAL_KSPLIT:
y_ptrs = (
y_ptr
+ ks * stride_yk
+ offs_m[:, None] * stride_ym
+ offs_n[None, :] * stride_yn
)
acc += tl.load(y_ptrs, mask=n_mask[None, :], other=0.0)
out = acc.to(out_ptr.type.element_ty)
out_ptrs = out_ptr + offs_m[:, None] * stride_om + offs_n[None, :] * stride_on
tl.store(out_ptrs, out, mask=n_mask[None, :])
def _safe_wrapper(a, b_shuffle, b_scale_sh):
a_q_raw, a_scale = dynamic_mxfp4_quant(a.contiguous())
a_scale_sh = _e8m0_shuffle_safe(a_scale)
return aiter.gemm_a4w4(
a_q_raw.view(_FP4X2), b_shuffle,
a_scale_sh.view(_FP8_E8M0), b_scale_sh,
dtype=_BF16, bpreshuffle=True,
)
def _get_route(m, n, k):
key = (m, n, k)
if key in _FUSED_CONFIGS:
return ("fused", _FUSED_CONFIGS[key])
if key in _ASM_CONFIGS:
return ("asm", _ASM_CONFIGS[key])
pair = (n, k)
if pair == (2112, 7168):
if m < 16:
return ("fused", _FUSED_K7168_KSPLIT14)
return ("fused", _FUSED_CONFIGS[(16, 2112, 7168)])
if pair == (3072, 1536):
if m <= 16:
return ("fused", _FUSED_K1536_SPLIT)
return ("fused", _FUSED_K1536_NOSPLIT)
if pair == (2880, 512):
if m < 16:
return ("fused", _FUSED_CONFIGS[(4, 2880, 512)])
if m < 128:
return ("fused", _FUSED_CONFIGS[(32, 2880, 512)])
return ("fused", _FUSED_BM32_K512)
if pair == (4096, 512):
return ("fused", _FUSED_CONFIGS[(32, 4096, 512)])
if pair == (7168, 2048):
return ("fused", _FUSED_K2048_NOSPLIT)
return None
def _make_asm_handler(m, k, n, device, splitk, kernel_name=_KERNEL_32X128):
padded_m = ((m + 31) >> 5) << 5
x_fp4 = torch.empty((padded_m, k >> 1), dtype=torch.uint8, device=device)
sN = (k + _QUANT_BLOCK - 1) // _QUANT_BLOCK
sN_pad = ((sN + 7) >> 3) << 3
sM_pad = ((m + 255) >> 8) << 8
scale = torch.empty((sM_pad, sN_pad), dtype=torch.uint8, device=device)
out = torch.empty_strided(
(padded_m, n), (n + 32, 1), dtype=_BF16, device=device,
)
x_fp4_view = x_fp4.view(_FP4X2)
scale_view = scale.view(_FP8_E8M0)
q_grid = ((m + _QUANT_TILE - 1) // _QUANT_TILE, sN_pad)
sfp4_0 = x_fp4.stride(0)
sfp4_1 = x_fp4.stride(1)
ssc_0 = scale.stride(0)
ssc_1 = scale.stride(1)
out_slice = out[:m]
quant_launch = _quant_kernel_asm_layout[q_grid]
def handler(a, b_shuffle, b_scale_sh):
quant_launch(
a, x_fp4, scale, k, 1, sfp4_0, sfp4_1, ssc_0, ssc_1,
M=m, N=k, scaleN=sN,
scaleM_pad=sM_pad, scaleN_pad=sN_pad,
BLOCK_SIZE=_QUANT_TILE, MXFP4_QUANT_BLOCK_SIZE=_QUANT_BLOCK, SHUFFLE=True,
)
gemm_a4w4_asm(
x_fp4_view, b_shuffle, scale_view, b_scale_sh,
out, kernel_name, bpreshuffle=True, log2_k_split=splitk,
)
return out_slice
return handler
def _make_fused_handler(m, n, k, device, raw_config, b_scale_sh):
config = dict(raw_config)
K_kernel = k // 2
if config["NUM_KSPLIT"] > 1:
SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT = get_splitk(
K_kernel, config["BLOCK_SIZE_K"], config["NUM_KSPLIT"]
)
config["SPLITK_BLOCK_SIZE"] = SPLITK_BLOCK_SIZE
config["BLOCK_SIZE_K"] = BLOCK_SIZE_K
config["NUM_KSPLIT"] = NUM_KSPLIT
if config["BLOCK_SIZE_K"] >= 2 * K_kernel:
config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * K_kernel)
config["SPLITK_BLOCK_SIZE"] = 2 * K_kernel
config["NUM_KSPLIT"] = 1
config["BLOCK_SIZE_N"] = max(config["BLOCK_SIZE_N"], 32)
if config["NUM_KSPLIT"] <= 1:
config["SPLITK_BLOCK_SIZE"] = 2 * K_kernel
has_splitk = config["NUM_KSPLIT"] > 1
out = torch.empty((m, n), dtype=_BF16, device=device)
y_pp = None
if has_splitk:
y_pp = torch.empty((config["NUM_KSPLIT"], m, n), dtype=torch.float32, device=device)
BM = config["BLOCK_SIZE_M"]
BN = config["BLOCK_SIZE_N"]
total_tiles = triton.cdiv(m, BM) * triton.cdiv(n, BN)
grid = (config["NUM_KSPLIT"] * total_tiles,)
exact_large_route = "m64" if (m, n, k) == (64, 7168, 2048) else ("m256" if (m, n, k) == (256, 3072, 1536) else "")
stride_ck = 0 if y_pp is None else y_pp.stride(0)
stride_cm = out.stride(0) if y_pp is None else y_pp.stride(1)
stride_cn = out.stride(1) if y_pp is None else y_pp.stride(2)
c_ptr = y_pp if has_splitk else out
a_s0 = k
w_s0 = k // 2 * 16
bss_1 = b_scale_sh.size(1)
ws_s0 = bss_1 * 32
use_exact_m16 = has_splitk and (m, n, k) == (16, 2112, 7168)
if use_exact_m16:
fused_launch = _gemm_exact_m16_xcd_kernel[grid]
elif has_splitk:
fused_launch = _gemm_optimized_xcd_kernel[grid]
else:
if exact_large_route == "m64":
grid = (224,)
fused_launch = _enhanced_preshuffle_exact_m64_kernel[grid]
elif exact_large_route == "m256":
grid = (384,)
fused_launch = _enhanced_preshuffle_exact_m256_kernel[grid]
else:
fused_launch = _enhanced_preshuffle_kernel[grid]
if has_splitk:
ACTUAL_KSPLIT = triton.cdiv(K_kernel, config["SPLITK_BLOCK_SIZE"] // 2)
MAX_KSPLIT = triton.next_power_of_2(config["NUM_KSPLIT"])
use_specialized_reduce = (m, n, k) == (16, 2112, 7168) and MAX_KSPLIT <= 8
if use_specialized_reduce:
reduce_grid = (triton.cdiv(n, 64),)
else:
reduce_grid = (triton.cdiv(m, 16), triton.cdiv(n, 64))
reduce_args = (
m, n, y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
out.stride(0), out.stride(1), 16, 64, ACTUAL_KSPLIT, MAX_KSPLIT,
)
reduce_launch = _gemm_afp4wfp4_reduce_kernel[reduce_grid]
if has_splitk:
def handler(a, b_shuffle, b_scale_sh):
w = b_shuffle.view(torch.uint8)
ws = b_scale_sh.view(torch.uint8)
fused_launch(a, w, c_ptr, ws, m, n, K_kernel, a_s0, 1, w_s0, 1,
stride_ck, stride_cm, stride_cn, ws_s0, 1, PREQUANT=True, **config)
if use_specialized_reduce:
_reduce_m16_splitk7_kernel[reduce_grid](
y_pp, out, n,
y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
out.stride(0), out.stride(1),
ACTUAL_KSPLIT=ACTUAL_KSPLIT, MAX_KSPLIT=MAX_KSPLIT, BLOCK_N=64,
num_warps=4, num_stages=1,
)
else:
reduce_launch(y_pp, out, *reduce_args)
return out
else:
def handler(a, b_shuffle, b_scale_sh):
w = b_shuffle.view(torch.uint8)
ws = b_scale_sh.view(torch.uint8)
fused_launch(a, w, c_ptr, ws, m, n, K_kernel, a_s0, 1, w_s0, 1,
stride_ck, stride_cm, stride_cn, ws_s0, 1, PREQUANT=True, **config)
return out
return handler
_HANDLERS = {}
_last_key = None
_last_handler = None
@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
global _last_key, _last_handler
a = data[0]
m = a.size(0)
k = a.size(1)
n = data[1].size(0)
key = (m, n, k)
if key == _last_key:
return _last_handler(a, data[3], data[4])
b_shuffle = data[3]
b_scale_sh = data[4]
if key not in _HANDLERS:
route = _get_route(m, n, k)
if route is None:
_HANDLERS[key] = lambda a, bs, bss: _safe_wrapper(a, bs, bss)
elif route[0] == "asm":
splitk, kname = route[1]
_HANDLERS[key] = _make_asm_handler(m, k, n, a.device, splitk, kname)
else:
_HANDLERS[key] = _make_fused_handler(m, n, k, a.device, route[1], b_scale_sh)
handler = _HANDLERS[key]
_last_key = key
_last_handler = handler
return handler(a, b_shuffle, b_scale_sh)
scrolls · 919 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 534310.
#!POPCORN leaderboard amd-mxfp4-mm#!POPCORN gpu MI355X+ """+ v1122_cand0873: exact M=16 XCD + M=16 waves=0 + M=256 stages=2 + M=32 stages=3.++ Key changes from cand_0902 base:+ - exact M=16 XCD kernel (hardcoded KSPLIT=7, NUM_K_ITER=2, pid decomposition)+ - M=16 waves_per_eu=0 (auto, was 2)+ - M=256 nosplit stages=2 (was 1)+ - M=32 shapes stages=3 (was 4)+ - duplicate exact M=16 kernel removed (faster JIT)++ Measured benchmark result:+ - ~7.40 us geomean (non-ranked)+ - [5.98, 8.55, 6.16, 6.00, 9.74, 9.81] us+ """from __future__ import annotations- """v423 plus fused routing for the hidden 3072x1536 M=64 family."""+ import os+ os.environ["TRITON_HIP_USE_BLOCK_PINGPONG"] = "0"import torchimport triton⋯ 2 unchanged linesfrom aiter import dtypesfrom aiter.ops.gemm_op_a4w4 import gemm_a4w4_asmfrom aiter.ops.triton.quant import dynamic_mxfp4_quant- from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op as _mxfp4_quant_op_even-+ from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op as _mxfp4_quant_op_sw+ from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid, remap_xcdfrom task import input_t, output_t++ # ---- Hardware FP4 conversion (prescale + scale=1.0 via VGPR) ----++ @triton.jit+ def _hw_fp4_convert_pair(val0, val1, scale):+ """Convert 2 FP32 values to packed FP4 byte using hardware instruction.+ CK-style: no v_mov_b32, "=v" output constraint, VGPR scale.+ """+ return tl.inline_asm_elementwise(+ "v_cvt_scalef32_pk_fp4_f32 $0, $1, $2, $3",+ "=v,v,v,v",+ [val0, val1, scale],+ dtype=tl.int32,+ is_pure=True,+ pack=1,+ )+++ @triton.jit+ def _mxfp4_quant_op_hw(+ x,+ BLOCK_SIZE_N,+ BLOCK_SIZE_M,+ MXFP4_QUANT_BLOCK_SIZE,+ ):+ """Hybrid quant: bitwise exponent extraction (no log2/floor) + exp2 prescale.+ Saves 2 SFU ops vs original hw quant while keeping low register pressure."""+ NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE+ x = x.to(tl.float32).reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)++ # Step 1: Bitwise exponent extraction (avoids log2 + floor)+ amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)+ amax_bits = amax.to(tl.int32, bitcast=True)+ amax_bits = (amax_bits + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000+ exponent_biased = ((amax_bits >> 23) & 0xFF).to(tl.int32)+ # scale_e8m0_unbiased = exponent_biased - 127 - 2 = exponent_biased - 129+ scale_e8m0_unbiased = exponent_biased - 129+ # tl.clamp doesn't support int32, use tl.where instead+ scale_e8m0_unbiased = tl.where(scale_e8m0_unbiased < -127, -127,+ tl.where(scale_e8m0_unbiased > 127, 127, scale_e8m0_unbiased))+ bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127++ # Step 2: Prescale via exp2 (keeps low register pressure, avoids bitwise prescale)+ prescale = tl.exp2((-scale_e8m0_unbiased).to(tl.float32))+ x = x * tl.broadcast_to(prescale, x.shape)++ # Step 3: Convert prescaled values to FP4 using hw instruction with scale=1.0+ x_pairs = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2)+ evens, odds = tl.split(x_pairs)+ ones = tl.full(evens.shape, 1.0, dtype=tl.float32)+ packed_i32 = _hw_fp4_convert_pair(evens, odds, ones)+ x_fp4 = (packed_i32 & 0xFF).to(tl.uint8)++ x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)+ return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)+++ @triton.jit+ def _mxfp4_quant_op_hw_bitwise(+ x,+ BLOCK_SIZE_N,+ BLOCK_SIZE_M,+ MXFP4_QUANT_BLOCK_SIZE,+ ):+ """Full bitwise quant — avoids log2, floor, AND exp2 via IEEE 754 manipulation."""+ NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE+ x = x.to(tl.float32).reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)++ # Step 1: Compute per-block scale via bitwise exponent extraction+ amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)+ amax_bits = amax.to(tl.int32, bitcast=True)+ amax_bits = (amax_bits + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000+ # Extract IEEE 754 biased exponent+ exponent_biased = ((amax_bits >> 23) & 0xFF).to(tl.int32)+ # bs_e8m0 = exponent_biased - 2, clamped to [0, 254]+ bs_e8m0_i32 = tl.where(exponent_biased < 2, 0, tl.where(exponent_biased > 256, 254, exponent_biased - 2))+ bs_e8m0 = bs_e8m0_i32.to(tl.uint8)++ # Step 2: Construct prescale as 2^(-scale_e8m0_unbiased) via bit manipulation+ # prescale_exponent_biased = 127 - scale_e8m0_unbiased = 127 - (exponent_biased - 129) = 256 - exponent_biased+ # After clamping: prescale_exponent_biased = 127 - (bs_e8m0_i32 - 127) = 254 - bs_e8m0_i32+ prescale_exp = 254 - bs_e8m0_i32+ prescale_bits = (prescale_exp << 23)+ prescale = prescale_bits.to(tl.float32, bitcast=True)+ x = x * tl.broadcast_to(prescale, x.shape)++ # Step 3: HW convert with scale=1.0+ x_pairs = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2)+ evens, odds = tl.split(x_pairs)+ ones = tl.full(evens.shape, 1.0, dtype=tl.float32)+ packed_i32 = _hw_fp4_convert_pair(evens, odds, ones)+ x_fp4 = (packed_i32 & 0xFF).to(tl.uint8)++ x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)+ return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)+++ # Monkey-patch BITWISE quant for preshuffle kernel path (faster for M=64, M=256)import aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 as _kernel_module+ _kernel_module._mxfp4_quant_op = _mxfp4_quant_op_hw_bitwise- _kernel_module._mxfp4_quant_op = _mxfp4_quant_op_even+ from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import _gemm_afp4wfp4_reduce_kernel+ from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk- from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle+ # ---- Enhanced preshuffle kernel: +XCD remap, +.wt store, +accumulator arg ----+ @triton.heuristics({+ "EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)+ and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)+ and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),+ "GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])+ * triton.cdiv(args["N"], args["BLOCK_SIZE_N"]),+ })+ @triton.jit+ def _enhanced_preshuffle_kernel(+ a_ptr, b_ptr, c_ptr, b_scales_ptr,+ M, N, K,+ stride_am, stride_ak, stride_bn, stride_bk,+ stride_ck, stride_cm, stride_cn, stride_bsn, stride_bsk,+ BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,+ BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr,+ NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,+ EVEN_K: tl.constexpr,+ num_warps: tl.constexpr, num_stages: tl.constexpr,+ waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,+ GRID_MN: tl.constexpr, PREQUANT: tl.constexpr, cache_modifier: tl.constexpr,+ ):+ 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_bsk > 0)+ tl.assume(stride_bsn > 0)+ pid_unified = tl.program_id(axis=0)+ pid_k = pid_unified % NUM_KSPLIT+ pid = pid_unified // NUM_KSPLIT+ # XCD remap ONLY for small grids ≤90 blocks (K=512 cases)+ # Hurts M=64 (224 blocks) and M=256 (384 blocks)+ # no remap+ num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)+ num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)+ if NUM_KSPLIT == 1:+ pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)+ else:+ pid_m = pid // num_pid_n+ pid_n = pid % num_pid_n+ SCALE_GROUP_SIZE: tl.constexpr = 32+ if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:+ num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)+ offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)+ offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16+ offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M+ a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak)+ offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)+ offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr+ offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N+ b_ptrs = b_ptr + (offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk)+ offs_bsn = (pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, (BLOCK_SIZE_N // 32))) % N+ offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)+ b_scale_ptrs = b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk+ accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)+ for k in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):+ b_scales = (tl.load(b_scale_ptrs, cache_modifier=cache_modifier)+ .reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8, 4, 16, 2, 2, 1)+ .permute(0, 5, 3, 1, 4, 2, 6)+ .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE))+ if EVEN_K:+ a_bf16 = tl.load(a_ptrs, cache_modifier=".ca")+ b = tl.load(b_ptrs, cache_modifier=cache_modifier)+ b = (b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)+ .permute(0, 1, 4, 2, 3, 5)+ .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)+ .trans(1, 0))+ if PREQUANT:+ a, a_scales = _mxfp4_quant_op_hw_bitwise(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)+ # CHANGE 2: Pass accumulator as arg (enables HW FMA fusion)+ accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)+ a_ptrs += BLOCK_SIZE_K * stride_ak+ b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk+ b_scale_ptrs += BLOCK_SIZE_K * stride_bsk+ c = accumulator.to(c_ptr.type.element_ty)+ offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)+ offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)+ c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :] + pid_k * stride_ck+ c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)+ if GRID_MN <= 224:+ tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")+ else:+ tl.store(c_ptrs, c, mask=c_mask)++ @triton.heuristics({+ "EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0),+ "GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])+ })+ @triton.jit+ def _enhanced_preshuffle_exact_m64_kernel(+ a_ptr, b_ptr, c_ptr, b_scales_ptr,+ M, N, K,+ stride_am, stride_ak, stride_bn, stride_bk,+ stride_ck, stride_cm, stride_cn, stride_bsn, stride_bsk,+ BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,+ BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr,+ NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,+ EVEN_K: tl.constexpr,+ num_warps: tl.constexpr, num_stages: tl.constexpr,+ waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,+ GRID_MN: tl.constexpr, PREQUANT: tl.constexpr, cache_modifier: tl.constexpr,+ ):+ pid = tl.program_id(axis=0)+ pid_m = pid // 56+ pid_n = pid % 56+ SCALE_GROUP_SIZE: tl.constexpr = 32+ NUM_K_ITER: tl.constexpr = 4+ offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)+ offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)+ a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_k_bf16[None, :] * stride_ak+ offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)+ offs_bn = pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)+ b_ptrs = b_ptr + offs_bn[:, None] * stride_bn + offs_k_shuffle_arr[None, :] * stride_bk+ offs_bsn = pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)+ offs_ks = tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)+ b_scale_ptrs = b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk+ accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)+ for _ in range(NUM_K_ITER):+ b_scales = (+ tl.load(b_scale_ptrs, cache_modifier=cache_modifier)+ .reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8, 4, 16, 2, 2, 1)+ .permute(0, 5, 3, 1, 4, 2, 6)+ .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)+ )+ a_bf16 = tl.load(a_ptrs, cache_modifier=".ca")+ b = tl.load(b_ptrs, cache_modifier=cache_modifier)+ b = (+ b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)+ .permute(0, 1, 4, 2, 3, 5)+ .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)+ .trans(1, 0)+ )+ if PREQUANT:+ a, a_scales = _mxfp4_quant_op_hw_bitwise(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)+ accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)+ a_ptrs += BLOCK_SIZE_K * stride_ak+ b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk+ b_scale_ptrs += BLOCK_SIZE_K * stride_bsk+ c = accumulator.to(c_ptr.type.element_ty)+ offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)+ offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)+ c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :]+ tl.store(c_ptrs, c, cache_modifier=".wt")+++ @triton.heuristics({+ "EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0),+ "GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])+ })+ @triton.jit+ def _enhanced_preshuffle_exact_m256_kernel(+ a_ptr, b_ptr, c_ptr, b_scales_ptr,+ M, N, K,+ stride_am, stride_ak, stride_bn, stride_bk,+ stride_ck, stride_cm, stride_cn, stride_bsn, stride_bsk,+ BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,+ BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr,+ NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,+ EVEN_K: tl.constexpr,+ num_warps: tl.constexpr, num_stages: tl.constexpr,+ waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,+ GRID_MN: tl.constexpr, PREQUANT: tl.constexpr, cache_modifier: tl.constexpr,+ ):+ pid = tl.program_id(axis=0)+ pid_m = pid // 24+ pid_n = pid % 24+ SCALE_GROUP_SIZE: tl.constexpr = 32+ NUM_K_ITER: tl.constexpr = 6+ offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)+ offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)+ a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_k_bf16[None, :] * stride_ak+ offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)+ offs_bn = pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)+ b_ptrs = b_ptr + offs_bn[:, None] * stride_bn + offs_k_shuffle_arr[None, :] * stride_bk+ offs_bsn = pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)+ offs_ks = tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)+ b_scale_ptrs = b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk+ accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)+ for _ in range(NUM_K_ITER):+ b_scales = (+ tl.load(b_scale_ptrs, cache_modifier=cache_modifier)+ .reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8, 4, 16, 2, 2, 1)+ .permute(0, 5, 3, 1, 4, 2, 6)+ .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)+ )+ a_bf16 = tl.load(a_ptrs, cache_modifier=".ca")+ b = tl.load(b_ptrs, cache_modifier=cache_modifier)+ b = (+ b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)+ .permute(0, 1, 4, 2, 3, 5)+ .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)+ .trans(1, 0)+ )+ if PREQUANT:+ a, a_scales = _mxfp4_quant_op_hw_bitwise(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)+ accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)+ a_ptrs += BLOCK_SIZE_K * stride_ak+ b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk+ b_scale_ptrs += BLOCK_SIZE_K * stride_bsk+ c = accumulator.to(c_ptr.type.element_ty)+ offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)+ offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)+ c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :]+ tl.store(c_ptrs, c)+++ # ---- Exact M=16 XCD kernel (hardcoded for M=16, N=2112, KSPLIT=7) ----+ @triton.heuristics({"EVEN_K": lambda args: True, "GRID_MN": lambda args: 34})+ @triton.jit+ def _gemm_exact_m16_xcd_kernel(+ a_ptr, b_ptr, c_ptr, b_scales_ptr,+ M, N, K,+ stride_am, stride_ak, stride_bn, stride_bk,+ stride_ck, stride_cm, stride_cn,+ stride_bsn, stride_bsk,+ BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,+ BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr,+ NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,+ EVEN_K: tl.constexpr,+ num_warps: tl.constexpr, num_stages: tl.constexpr,+ waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,+ GRID_MN: tl.constexpr,+ PREQUANT: tl.constexpr, cache_modifier: tl.constexpr,+ ):+ 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_bsk > 0)+ tl.assume(stride_bsn > 0)+ pid_linear = tl.program_id(axis=0)+ pid_k = pid_linear % 7+ pid = remap_xcd(pid_linear // 7, 34, NUM_XCDS=8)+ pid_m = pid % 2+ pid_n = pid // 2+ SCALE_GROUP_SIZE: tl.constexpr = 32+ NUM_K_ITER: tl.constexpr = 2+ offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)+ offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16+ offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)+ a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak+ offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)+ offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr+ offs_bn = pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)+ b_ptrs = b_ptr + offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk+ offs_bsn = pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)+ offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)+ b_scale_ptrs = b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk+ accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)+ for _ in range(NUM_K_ITER):+ b_scales = (+ tl.load(b_scale_ptrs, cache_modifier=cache_modifier)+ .reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8, 4, 16, 2, 2, 1)+ .permute(0, 5, 3, 1, 4, 2, 6)+ .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)+ )+ a_bf16 = tl.load(a_ptrs, cache_modifier=".ca")+ b = tl.load(b_ptrs, cache_modifier=cache_modifier)+ b = (+ b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)+ .permute(0, 1, 4, 2, 3, 5)+ .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)+ .trans(1, 0)+ )+ a, a_scales = _mxfp4_quant_op_hw(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)+ accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)+ a_ptrs += BLOCK_SIZE_K * stride_ak+ b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk+ b_scale_ptrs += BLOCK_SIZE_K * stride_bsk+ c = accumulator.to(c_ptr.type.element_ty)+ offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)+ offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)+ c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :] + pid_k * stride_ck+ c_mask = offs_cn[None, :] < N+ tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")+_BF16 = dtypes.bf16_FP4X2 = dtypes.fp4x2_FP8_E8M0 = dtypes.fp8_e8m0_KERNEL_32X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"- _PUBLIC_SMALL = {+ _FUSED_CONFIGS = {(4, 2880, 512): {- "BLOCK_SIZE_M": 8,- "BLOCK_SIZE_N": 64,- "BLOCK_SIZE_K": 512,- "GROUP_SIZE_M": 1,- "num_warps": 4,- "num_stages": 2,- "waves_per_eu": 2,- "matrix_instr_nonkdim": 16,- "cache_modifier": ".cg",- "NUM_KSPLIT": 1,+ "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256,+ "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3,+ "waves_per_eu": 2, "matrix_instr_nonkdim": 16,+ "cache_modifier": ".cg", "NUM_KSPLIT": 1,},(16, 2112, 7168): {- "BLOCK_SIZE_M": 16,- "BLOCK_SIZE_N": 64,- "BLOCK_SIZE_K": 512,- "GROUP_SIZE_M": 1,- "num_warps": 4,- "num_stages": 2,- "waves_per_eu": 1,- "matrix_instr_nonkdim": 16,- "cache_modifier": ".cg",- "NUM_KSPLIT": 8,+ "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,+ "GROUP_SIZE_M": 2, "num_warps": 4, "num_stages": 2,+ "waves_per_eu": 0, "matrix_instr_nonkdim": 16,+ "cache_modifier": ".cg", "NUM_KSPLIT": 7,},(32, 4096, 512): {- "BLOCK_SIZE_M": 8,- "BLOCK_SIZE_N": 64,- "BLOCK_SIZE_K": 512,- "GROUP_SIZE_M": 1,- "num_warps": 4,- "num_stages": 3,- "waves_per_eu": 1,- "matrix_instr_nonkdim": 16,- "cache_modifier": ".cg",- "NUM_KSPLIT": 8,+ "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256,+ "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3,+ "waves_per_eu": 1, "matrix_instr_nonkdim": 16,+ "cache_modifier": ".cg", "NUM_KSPLIT": 1,},(32, 2880, 512): {- "BLOCK_SIZE_M": 8,- "BLOCK_SIZE_N": 64,- "BLOCK_SIZE_K": 512,- "GROUP_SIZE_M": 1,- "num_warps": 4,- "num_stages": 2,- "waves_per_eu": 1,- "matrix_instr_nonkdim": 16,- "cache_modifier": None,- "NUM_KSPLIT": 1,+ "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256,+ "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3,+ "waves_per_eu": 1, "matrix_instr_nonkdim": 16,+ "cache_modifier": None, "NUM_KSPLIT": 1,},+ # M=256 via fused preshuffle — v925: BM=16 + cache=".cg"+ (256, 3072, 1536): {+ "BLOCK_SIZE_M": 16, "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": 1,+ },}- _PUBLIC_LARGE = {- (64, 7168, 2048): 2,- (256, 3072, 1536): 1,+ _FUSED_K7168_KSPLIT14 = {+ "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,+ "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3,+ "waves_per_eu": 1, "matrix_instr_nonkdim": 16,+ "cache_modifier": ".cg", "NUM_KSPLIT": 14,}- _PUBLIC_TEST_SMALL = {- (8, 2112, 7168): {- "BLOCK_SIZE_M": 8,- "BLOCK_SIZE_N": 128,- "BLOCK_SIZE_K": 512,- "GROUP_SIZE_M": 1,- "num_warps": 4,- "num_stages": 1,- "waves_per_eu": 1,- "matrix_instr_nonkdim": 16,- "cache_modifier": ".cg",- "NUM_KSPLIT": 14,- },- (16, 3072, 1536): {- "BLOCK_SIZE_M": 16,- "BLOCK_SIZE_N": 128,- "BLOCK_SIZE_K": 512,- "GROUP_SIZE_M": 1,- "num_warps": 4,- "num_stages": 1,- "waves_per_eu": 1,- "matrix_instr_nonkdim": 16,- "cache_modifier": ".cg",- "NUM_KSPLIT": 3,- },+ _FUSED_K1536_SPLIT = {+ "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,+ "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3,+ "waves_per_eu": 1, "matrix_instr_nonkdim": 16,+ "cache_modifier": ".cg", "NUM_KSPLIT": 3,}+ _FUSED_K1536_NOSPLIT = {+ "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256,+ "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,+ "waves_per_eu": 1, "matrix_instr_nonkdim": 16,+ "cache_modifier": ".cg", "NUM_KSPLIT": 1,+ }++ _FUSED_BM32_K512 = {+ "BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,+ "GROUP_SIZE_M": 1, "num_warps": 8, "num_stages": 3,+ "waves_per_eu": 1, "matrix_instr_nonkdim": 16,+ "cache_modifier": ".cg", "NUM_KSPLIT": 1,+ }++ _FUSED_K2048_NOSPLIT = {+ "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,+ "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,+ "waves_per_eu": 1, "matrix_instr_nonkdim": 16,+ "cache_modifier": ".cg", "NUM_KSPLIT": 1,+ }++ _ASM_CONFIGS = {+ # Removed M=256 from ASM — use fused preshuffle instead+ }+_QUANT_BLOCK = 32_QUANT_TILE = 128- _OUT_PAD_BF16 = 32- _BUFS = {}+ # ---- Standalone quant kernel for ASM path (SOFTWARE quant) ----+@triton.jit- def _dynamic_mxfp4_quant_kernel_even_asm_layout(- x_ptr,- x_fp4_ptr,- bs_ptr,- stride_x_m,- stride_x_n,- stride_x_fp4_m,- stride_x_fp4_n,- stride_bs_m,- stride_bs_n,- M: tl.constexpr,- N: tl.constexpr,- scaleN: tl.constexpr,- scaleM_pad: tl.constexpr,- scaleN_pad: tl.constexpr,- BLOCK_SIZE: tl.constexpr,- MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,+ def _quant_kernel_asm_layout(+ x_ptr, x_fp4_ptr, bs_ptr,+ stride_x_m, stride_x_n,+ stride_x_fp4_m, stride_x_fp4_n,+ stride_bs_m, stride_bs_n,+ M: tl.constexpr, N: tl.constexpr,+ scaleN: tl.constexpr, scaleM_pad: tl.constexpr, scaleN_pad: tl.constexpr,+ BLOCK_SIZE: tl.constexpr, MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,SHUFFLE: tl.constexpr,):pid_m = tl.program_id(0)pid_n = tl.program_id(1)-stride_x_m = tl.cast(stride_x_m, tl.int64)stride_x_n = tl.cast(stride_x_n, tl.int64)stride_x_fp4_m = tl.cast(stride_x_fp4_m, tl.int64)stride_x_fp4_n = tl.cast(stride_x_fp4_n, tl.int64)-x_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)x_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE)x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_nx_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]x = tl.load(x_ptr + x_offs, mask=x_mask).to(tl.float32)-- out_tensor, bs_e8m0 = _mxfp4_quant_op_even(- x,- MXFP4_QUANT_BLOCK_SIZE,- BLOCK_SIZE,- MXFP4_QUANT_BLOCK_SIZE,+ # Use SOFTWARE quant for ASM path+ out_tensor, bs_e8m0 = _mxfp4_quant_op_sw(+ x, MXFP4_QUANT_BLOCK_SIZE, BLOCK_SIZE, MXFP4_QUANT_BLOCK_SIZE,)-out_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)- out_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE // 2 + tl.arange(- 0, MXFP4_QUANT_BLOCK_SIZE // 2- )- out_offs = (- out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n- )+ out_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE // 2 + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE // 2)+ out_offs = out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_nout_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)-bs_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)bs_offs_n = pid_n-if SHUFFLE:bs_offs_0 = bs_offs_m[:, None] // 32bs_offs_1 = bs_offs_m[:, None] % 32⋯ 4 unchanged linesbs_offs_5 = bs_offs_4 % 4bs_offs_4 = bs_offs_4 // 4bs_offs = (- bs_offs_1- + bs_offs_4 * 2- + bs_offs_2 * 4- + bs_offs_5 * 64- + bs_offs_3 * 256+ bs_offs_1 + bs_offs_4 * 2 + bs_offs_2 * 4+ + bs_offs_5 * 64 + bs_offs_3 * 256+ bs_offs_0 * 32 * scaleN)bs_mask1 = (bs_offs_m < M)[:, None] & (bs_offs_n < scaleN)[None, :]⋯ 6 unchanged linestl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask)++ # ---- Optimized XCD kernel (fused, uses hw quant) ----++ @triton.heuristics(+ {+ "EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)+ and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)+ and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),+ "GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])+ * triton.cdiv(args["N"], args["BLOCK_SIZE_N"]),+ }+ )+ @triton.jit+ def _gemm_optimized_xcd_kernel(+ a_ptr, b_ptr, c_ptr, b_scales_ptr,+ M, N, K,+ stride_am, stride_ak, stride_bn, stride_bk,+ stride_ck, stride_cm, stride_cn,+ stride_bsn, stride_bsk,+ BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,+ BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr,+ NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,+ EVEN_K: tl.constexpr,+ num_warps: tl.constexpr, num_stages: tl.constexpr,+ waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,+ GRID_MN: tl.constexpr,+ PREQUANT: tl.constexpr, cache_modifier: tl.constexpr,+ ):+ 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_bsk > 0)+ tl.assume(stride_bsn > 0)++ pid_linear = tl.program_id(axis=0)+ pid_k = pid_linear % NUM_KSPLIT+ pid = remap_xcd(pid_linear // NUM_KSPLIT, GRID_MN, NUM_XCDS=8)+ num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)+ num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)+ pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)++ SCALE_GROUP_SIZE: tl.constexpr = 32++ if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:+ num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)+ offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)+ offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16+ offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M+ a_ptrs = a_ptr + (+ offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak+ )+ offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)+ offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr+ offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N+ b_ptrs = b_ptr + (+ offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk+ )+ offs_bsn = (+ pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, (BLOCK_SIZE_N // 32))+ ) % N+ offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(+ 0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32+ )+ b_scale_ptrs = (+ b_scales_ptr+ + offs_bsn[:, None] * stride_bsn+ + offs_ks[None, :] * stride_bsk+ )+ accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)+ for _ in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):+ b_scales = (+ tl.load(b_scale_ptrs, cache_modifier=cache_modifier)+ .reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8, 4, 16, 2, 2, 1)+ .permute(0, 5, 3, 1, 4, 2, 6)+ .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)+ )+ if EVEN_K:+ a_bf16 = tl.load(a_ptrs, cache_modifier=".ca")+ b = tl.load(b_ptrs, cache_modifier=cache_modifier)+ b = (+ b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)+ .permute(0, 1, 4, 2, 3, 5)+ .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)+ .trans(1, 0)+ )+ if PREQUANT:+ a, a_scales = _mxfp4_quant_op_hw(+ a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32+ )+ accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)+ a_ptrs += BLOCK_SIZE_K * stride_ak+ b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk+ b_scale_ptrs += BLOCK_SIZE_K * stride_bsk++ c = accumulator.to(c_ptr.type.element_ty)+ offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)+ offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)+ c_ptrs = (+ c_ptr+ + stride_cm * offs_cm[:, None]+ + stride_cn * offs_cn[None, :]+ + pid_k * stride_ck+ )+ c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)+ if GRID_MN <= 224:+ tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")+ else:+ tl.store(c_ptrs, c, mask=c_mask)++def _e8m0_shuffle_safe(scale: torch.Tensor) -> torch.Tensor:m, n = scale.shapescale_padded = torch.empty(- ((m + 255) // 256) * 256,- ((n + 7) // 8) * 8,- dtype=scale.dtype,- device=scale.device,+ ((m + 255) // 256) * 256, ((n + 7) // 8) * 8,+ dtype=scale.dtype, device=scale.device,)scale_padded.fill_(0x7F)scale_padded[:m, :n] = scalesm, sn = scale_padded.shapereturn (scale_padded.view(sm // 32, 2, 16, sn // 8, 2, 4)- .permute(0, 3, 5, 2, 4, 1)- .contiguous()- .view(sm, sn)+ .permute(0, 3, 5, 2, 4, 1).contiguous().view(sm, sn))- def _safe_wrapper(a: torch.Tensor, b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor):+ @triton.jit+ def _reduce_m16_splitk7_kernel(+ y_ptr, out_ptr,+ N,+ stride_yk, stride_ym, stride_yn,+ stride_om, stride_on,+ ACTUAL_KSPLIT: tl.constexpr,+ MAX_KSPLIT: tl.constexpr,+ BLOCK_N: tl.constexpr,+ ):+ pid_n = tl.program_id(axis=0)+ offs_m = tl.arange(0, 16)+ offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)+ n_mask = offs_n < N+ acc = tl.zeros((16, BLOCK_N), dtype=tl.float32)++ for ks in range(MAX_KSPLIT):+ if ks < ACTUAL_KSPLIT:+ y_ptrs = (+ y_ptr+ + ks * stride_yk+ + offs_m[:, None] * stride_ym+ + offs_n[None, :] * stride_yn+ )+ acc += tl.load(y_ptrs, mask=n_mask[None, :], other=0.0)++ out = acc.to(out_ptr.type.element_ty)+ out_ptrs = out_ptr + offs_m[:, None] * stride_om + offs_n[None, :] * stride_on+ tl.store(out_ptrs, out, mask=n_mask[None, :])+++ def _safe_wrapper(a, b_shuffle, b_scale_sh):a_q_raw, a_scale = dynamic_mxfp4_quant(a.contiguous())a_scale_sh = _e8m0_shuffle_safe(a_scale)return aiter.gemm_a4w4(- a_q_raw.view(_FP4X2),- b_shuffle,- a_scale_sh.view(_FP8_E8M0),- b_scale_sh,- dtype=_BF16,- bpreshuffle=True,+ a_q_raw.view(_FP4X2), b_shuffle,+ a_scale_sh.view(_FP8_E8M0), b_scale_sh,+ dtype=_BF16, bpreshuffle=True,)- def _get_large_bufs(m: int, k: int, n: int, device):- x_fp4 = torch.empty((m, k >> 1), dtype=torch.uint8, device=device)- scale_n = (k + _QUANT_BLOCK - 1) // _QUANT_BLOCK- scale_n_pad = ((scale_n + 7) >> 3) << 3- scale_m_pad = ((m + 255) >> 8) << 8- scale = torch.empty((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=device)- padded_m = ((m + 31) >> 5) << 5- out = torch.empty_strided(- (padded_m, n),- (n + _OUT_PAD_BF16, 1),- dtype=_BF16,- device=device,- )- return x_fp4, scale, scale_n, scale_n_pad, scale_m_pad, out--- def _get_route(key):- m, n, k = key-- if key in _PUBLIC_SMALL:- return ("small", _PUBLIC_SMALL[key])- if key in _PUBLIC_LARGE:- return ("large", _PUBLIC_LARGE[key])-+ def _get_route(m, n, k):+ key = (m, n, k)+ if key in _FUSED_CONFIGS:+ return ("fused", _FUSED_CONFIGS[key])+ if key in _ASM_CONFIGS:+ return ("asm", _ASM_CONFIGS[key])pair = (n, k)if pair == (2112, 7168):if m < 16:- return ("small", _PUBLIC_TEST_SMALL[(8, 2112, 7168)])- return ("small", _PUBLIC_SMALL[(16, 2112, 7168)])+ return ("fused", _FUSED_K7168_KSPLIT14)+ return ("fused", _FUSED_CONFIGS[(16, 2112, 7168)])if pair == (3072, 1536):- if m < 128:- return ("small", _PUBLIC_TEST_SMALL[(16, 3072, 1536)])- return ("large", 1)+ if m <= 16:+ return ("fused", _FUSED_K1536_SPLIT)+ return ("fused", _FUSED_K1536_NOSPLIT)if pair == (2880, 512):if m < 16:- return ("small", _PUBLIC_SMALL[(4, 2880, 512)])- if m < 96:- return ("small", _PUBLIC_SMALL[(32, 2880, 512)])- return ("large", 2)+ return ("fused", _FUSED_CONFIGS[(4, 2880, 512)])+ if m < 128:+ return ("fused", _FUSED_CONFIGS[(32, 2880, 512)])+ return ("fused", _FUSED_BM32_K512)if pair == (4096, 512):- return ("small", _PUBLIC_SMALL[(32, 4096, 512)])+ return ("fused", _FUSED_CONFIGS[(32, 4096, 512)])if pair == (7168, 2048):- return ("large", 2)+ return ("fused", _FUSED_K2048_NOSPLIT)return None- @torch.inference_mode()- def custom_kernel(data: input_t) -> output_t:- a, b, _b_q, b_shuffle, b_scale_sh = data- m, k = a.shape- n = b.shape[0]- key = (int(m), int(n), int(k))- route = _get_route(key)+ def _make_asm_handler(m, k, n, device, splitk, kernel_name=_KERNEL_32X128):+ padded_m = ((m + 31) >> 5) << 5+ x_fp4 = torch.empty((padded_m, k >> 1), dtype=torch.uint8, device=device)+ sN = (k + _QUANT_BLOCK - 1) // _QUANT_BLOCK+ sN_pad = ((sN + 7) >> 3) << 3+ sM_pad = ((m + 255) >> 8) << 8+ scale = torch.empty((sM_pad, sN_pad), dtype=torch.uint8, device=device)+ out = torch.empty_strided(+ (padded_m, n), (n + 32, 1), dtype=_BF16, device=device,+ )+ x_fp4_view = x_fp4.view(_FP4X2)+ scale_view = scale.view(_FP8_E8M0)+ q_grid = ((m + _QUANT_TILE - 1) // _QUANT_TILE, sN_pad)+ sfp4_0 = x_fp4.stride(0)+ sfp4_1 = x_fp4.stride(1)+ ssc_0 = scale.stride(0)+ ssc_1 = scale.stride(1)+ out_slice = out[:m]+ quant_launch = _quant_kernel_asm_layout[q_grid]- if route is None:- return _safe_wrapper(a, b_shuffle, b_scale_sh)-- route_kind, route_value = route-- if route_kind == "large":- if key not in _BUFS:- _BUFS[key] = ("large", _get_large_bufs(m, k, n, a.device))- _, (x_fp4, scale, scale_n, scale_n_pad, scale_m_pad, out) = _BUFS[key]- grid = ((m + _QUANT_TILE - 1) // _QUANT_TILE, scale_n_pad)- _dynamic_mxfp4_quant_kernel_even_asm_layout[grid](- a,- x_fp4,- scale,- a.stride(0),- a.stride(1),- x_fp4.stride(0),- x_fp4.stride(1),- scale.stride(0),- scale.stride(1),- M=m,- N=k,- scaleN=scale_n,- scaleM_pad=scale_m_pad,- scaleN_pad=scale_n_pad,- BLOCK_SIZE=_QUANT_TILE,- MXFP4_QUANT_BLOCK_SIZE=_QUANT_BLOCK,- SHUFFLE=True,+ def handler(a, b_shuffle, b_scale_sh):+ quant_launch(+ a, x_fp4, scale, k, 1, sfp4_0, sfp4_1, ssc_0, ssc_1,+ M=m, N=k, scaleN=sN,+ scaleM_pad=sM_pad, scaleN_pad=sN_pad,+ BLOCK_SIZE=_QUANT_TILE, MXFP4_QUANT_BLOCK_SIZE=_QUANT_BLOCK, SHUFFLE=True,)gemm_a4w4_asm(- x_fp4.view(_FP4X2),- b_shuffle,- scale.view(_FP8_E8M0),- b_scale_sh,- out,- _KERNEL_32X128,- bpreshuffle=True,- log2_k_split=route_value,+ x_fp4_view, b_shuffle, scale_view, b_scale_sh,+ out, kernel_name, bpreshuffle=True, log2_k_split=splitk,)- return out[:m]+ return out_slice- if key not in _BUFS:- _BUFS[key] = ("small", torch.empty((m, n), dtype=_BF16, device=a.device))- _, out = _BUFS[key]+ return handler- w = b_shuffle.view(torch.uint8).reshape(n // 16, k // 2 * 16)- sm, sn = b_scale_sh.shape- w_scales = b_scale_sh.view(torch.uint8).reshape(sm // 32, sn * 32)- return gemm_a16wfp4_preshuffle(- a,- w,- w_scales,- prequant=True,- y=out,- config=route_value,- )++ def _make_fused_handler(m, n, k, device, raw_config, b_scale_sh):+ config = dict(raw_config)+ K_kernel = k // 2+ if config["NUM_KSPLIT"] > 1:+ SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT = get_splitk(+ K_kernel, config["BLOCK_SIZE_K"], config["NUM_KSPLIT"]+ )+ config["SPLITK_BLOCK_SIZE"] = SPLITK_BLOCK_SIZE+ config["BLOCK_SIZE_K"] = BLOCK_SIZE_K+ config["NUM_KSPLIT"] = NUM_KSPLIT+ if config["BLOCK_SIZE_K"] >= 2 * K_kernel:+ config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * K_kernel)+ config["SPLITK_BLOCK_SIZE"] = 2 * K_kernel+ config["NUM_KSPLIT"] = 1+ config["BLOCK_SIZE_N"] = max(config["BLOCK_SIZE_N"], 32)+ if config["NUM_KSPLIT"] <= 1:+ config["SPLITK_BLOCK_SIZE"] = 2 * K_kernel++ has_splitk = config["NUM_KSPLIT"] > 1+ out = torch.empty((m, n), dtype=_BF16, device=device)+ y_pp = None+ if has_splitk:+ y_pp = torch.empty((config["NUM_KSPLIT"], m, n), dtype=torch.float32, device=device)++ BM = config["BLOCK_SIZE_M"]+ BN = config["BLOCK_SIZE_N"]+ total_tiles = triton.cdiv(m, BM) * triton.cdiv(n, BN)+ grid = (config["NUM_KSPLIT"] * total_tiles,)+ exact_large_route = "m64" if (m, n, k) == (64, 7168, 2048) else ("m256" if (m, n, k) == (256, 3072, 1536) else "")+ stride_ck = 0 if y_pp is None else y_pp.stride(0)+ stride_cm = out.stride(0) if y_pp is None else y_pp.stride(1)+ stride_cn = out.stride(1) if y_pp is None else y_pp.stride(2)+ c_ptr = y_pp if has_splitk else out+ a_s0 = k+ w_s0 = k // 2 * 16+ bss_1 = b_scale_sh.size(1)+ ws_s0 = bss_1 * 32++ use_exact_m16 = has_splitk and (m, n, k) == (16, 2112, 7168)+ if use_exact_m16:+ fused_launch = _gemm_exact_m16_xcd_kernel[grid]+ elif has_splitk:+ fused_launch = _gemm_optimized_xcd_kernel[grid]+ else:+ if exact_large_route == "m64":+ grid = (224,)+ fused_launch = _enhanced_preshuffle_exact_m64_kernel[grid]+ elif exact_large_route == "m256":+ grid = (384,)+ fused_launch = _enhanced_preshuffle_exact_m256_kernel[grid]+ else:+ fused_launch = _enhanced_preshuffle_kernel[grid]++ if has_splitk:+ ACTUAL_KSPLIT = triton.cdiv(K_kernel, config["SPLITK_BLOCK_SIZE"] // 2)+ MAX_KSPLIT = triton.next_power_of_2(config["NUM_KSPLIT"])+ use_specialized_reduce = (m, n, k) == (16, 2112, 7168) and MAX_KSPLIT <= 8+ if use_specialized_reduce:+ reduce_grid = (triton.cdiv(n, 64),)+ else:+ reduce_grid = (triton.cdiv(m, 16), triton.cdiv(n, 64))+ reduce_args = (+ m, n, y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),+ out.stride(0), out.stride(1), 16, 64, ACTUAL_KSPLIT, MAX_KSPLIT,+ )+ reduce_launch = _gemm_afp4wfp4_reduce_kernel[reduce_grid]++ if has_splitk:+ def handler(a, b_shuffle, b_scale_sh):+ w = b_shuffle.view(torch.uint8)+ ws = b_scale_sh.view(torch.uint8)+ fused_launch(a, w, c_ptr, ws, m, n, K_kernel, a_s0, 1, w_s0, 1,+ stride_ck, stride_cm, stride_cn, ws_s0, 1, PREQUANT=True, **config)+ if use_specialized_reduce:+ _reduce_m16_splitk7_kernel[reduce_grid](+ y_pp, out, n,+ y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),+ out.stride(0), out.stride(1),+ ACTUAL_KSPLIT=ACTUAL_KSPLIT, MAX_KSPLIT=MAX_KSPLIT, BLOCK_N=64,+ num_warps=4, num_stages=1,+ )+ else:+ reduce_launch(y_pp, out, *reduce_args)+ return out+ else:+ def handler(a, b_shuffle, b_scale_sh):+ w = b_shuffle.view(torch.uint8)+ ws = b_scale_sh.view(torch.uint8)+ fused_launch(a, w, c_ptr, ws, m, n, K_kernel, a_s0, 1, w_s0, 1,+ stride_ck, stride_cm, stride_cn, ws_s0, 1, PREQUANT=True, **config)+ return out++ return handler+++ _HANDLERS = {}+ _last_key = None+ _last_handler = None+++ @torch.inference_mode()+ def custom_kernel(data: input_t) -> output_t:+ global _last_key, _last_handler+ a = data[0]+ m = a.size(0)+ k = a.size(1)+ n = data[1].size(0)+ key = (m, n, k)+ if key == _last_key:+ return _last_handler(a, data[3], data[4])+ b_shuffle = data[3]+ b_scale_sh = data[4]+ if key not in _HANDLERS:+ route = _get_route(m, n, k)+ if route is None:+ _HANDLERS[key] = lambda a, bs, bss: _safe_wrapper(a, bs, bss)+ elif route[0] == "asm":+ splitk, kname = route[1]+ _HANDLERS[key] = _make_asm_handler(m, k, n, a.device, splitk, kname)+ else:+ _HANDLERS[key] = _make_fused_handler(m, n, k, a.device, route[1], b_scale_sh)+ handler = _HANDLERS[key]+ _last_key = key+ _last_handler = handler+ return handler(a, b_shuffle, b_scale_sh)
scrolls · 1131 diff lines total
Best evidence level for this revision: reported
JSON