submission 708351
ihansel · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 338 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-708351?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:beeb3a2c0f9736b03fb7a60f33d5059c6b877f4817b3ca467d6493171366b514
license declaredunknown
license concludedunknown
authorsihansel
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
"""MXFP4 GEMM v129: Hybrid best-of-both.num-warps = 4
NUM_KSPLIT=NUM_KSPLIT, BLOCK_M=16, BLOCK_N=64, num_warps=4,split-k
splitK = 0stages = 2
num_warps=num_warps, num_stages=2, matrix_instr_nonkdim=16,tile-m = 16
NUM_KSPLIT=NUM_KSPLIT, BLOCK_M=16, BLOCK_N=64, num_warps=4,tile-n = 64
NUM_KSPLIT=NUM_KSPLIT, BLOCK_M=16, BLOCK_N=64, num_warps=4,Kernel source
submission.py338 lines
"""MXFP4 GEMM v129: Hybrid best-of-both.
M<=32: AITER quant fused into Triton GEMM (v128, 7.3-10.4µs)
M>=64: fused quant+shuffle + direct ASM (v126, 15.9-17.2µs)
"""
import torch
import triton
import triton.language as tl
import aiter # noqa: F401
from aiter import dtypes
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm, get_GEMM_config
from task import input_t, output_t
from reference import ref_kernel # noqa: F401
# ==================== FUSED QUANT+SHUFFLE KERNEL (for M>=64 ASM path) ====================
@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 _fused_quant_shuffle_kernel(
x_ptr, x_fp4_ptr, bs_ptr,
stride_x_m_in, stride_x_n_in, stride_x_fp4_m_in, stride_x_fp4_n_in,
M: tl.constexpr, N: tl.constexpr, scaleN_valid: tl.constexpr, scaleN_pad: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr, SCALING_MODE: tl.constexpr,
NUM_ITER: tl.constexpr, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
NUM_STAGES: tl.constexpr, EVEN_M_N: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
stride_x_m = tl.cast(stride_x_m_in, tl.int64)
stride_x_n = tl.cast(stride_x_n_in, tl.int64)
stride_x_fp4_m = tl.cast(stride_x_fp4_m_in, tl.int64)
stride_x_fp4_n = tl.cast(stride_x_fp4_n_in, tl.int64)
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
for iter in tl.static_range(NUM_ITER):
offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
offs_n = (pid_n * NUM_ITER + iter) * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
x_offs = offs_m[:, None] * stride_x_m + offs_n[None, :] * stride_x_n
if EVEN_M_N:
x = tl.load(x_ptr + x_offs).to(tl.float32)
else:
x_mask = (offs_m < M)[:, None] & (offs_n < N)[None, :]
x = tl.load(x_ptr + x_offs, mask=x_mask, other=0.0).to(tl.float32)
x_fp4, bs_e8m0 = _mxfp4_quant_op(x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE)
offs_fp4_n = (pid_n * NUM_ITER + iter) * (BLOCK_SIZE_N // 2) + tl.arange(0, BLOCK_SIZE_N // 2)
out_offs = offs_m[:, None] * stride_x_fp4_m + offs_fp4_n[None, :] * stride_x_fp4_n
if EVEN_M_N:
tl.store(x_fp4_ptr + out_offs, x_fp4)
else:
out_mask = (offs_m < M)[:, None] & (offs_fp4_n < (N // 2))[None, :]
tl.store(x_fp4_ptr + out_offs, x_fp4, mask=out_mask)
bs_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
bs_col_base = (pid_n * NUM_ITER + iter) * NUM_QUANT_BLOCKS
bs_offs_n = bs_col_base + tl.arange(0, NUM_QUANT_BLOCKS)
g0 = bs_offs_m // 32
rem32 = bs_offs_m % 32
g1 = rem32 // 16
g2 = rem32 % 16
g3 = bs_offs_n // 8
g4 = (bs_offs_n % 8) // 4
g5 = bs_offs_n % 4
shuffled_offs = (g1[:, None] + g4[None, :] * 2 + g2[:, None] * 4
+ g5[None, :] * 64 + g3[None, :] * 256
+ g0[:, None] * (scaleN_pad * 32))
bs_mask_valid = (bs_offs_m < M)[:, None] & (bs_offs_n < scaleN_valid)[None, :]
bs_e8m0_padded = tl.where(bs_mask_valid, bs_e8m0, 127)
sm_pad = ((M + 255) // 256) * 256
bs_mask_write = (bs_offs_m < sm_pad)[:, None] & (bs_offs_n < scaleN_pad)[None, :]
tl.store(bs_ptr + shuffled_offs, bs_e8m0_padded, mask=bs_mask_write)
@triton.jit
def _aiter_fused_quant_gemm_kernel(
a_bf16_ptr, b_ptr, c_ptr, b_scales_ptr,
M, N, K,
stride_am, stride_ak, stride_bk, stride_bn,
stride_cs, stride_cm, stride_cn,
stride_bsn, stride_bsk,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
NUM_KSPLIT: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
):
SCALE_GROUP_SIZE: tl.constexpr = MXFP4_QUANT_BLOCK_SIZE
pid = tl.program_id(axis=0)
num_pid_m = tl.cdiv(M, BLOCK_M)
num_pid_n = tl.cdiv(N, BLOCK_N)
num_pid_mn = num_pid_m * num_pid_n
pid_k = pid // num_pid_mn
pid_mn = pid % num_pid_mn
pid_m = pid_mn // num_pid_n
pid_n = pid_mn % num_pid_n
total_k_iters = tl.cdiv(K, BLOCK_K)
k_iters_per_split = tl.cdiv(total_k_iters, NUM_KSPLIT)
k_start = pid_k * k_iters_per_split
k_end = min((pid_k + 1) * k_iters_per_split, total_k_iters)
offs_am = (pid_m * BLOCK_M + tl.arange(0, BLOCK_M)) % M
offs_k_bf16 = tl.arange(0, BLOCK_K)
a_bf16_ptrs = a_bf16_ptr + (offs_am[:, None] * stride_am + offs_k_bf16[None, :] * stride_ak)
a_bf16_ptrs += k_start * BLOCK_K * stride_ak
offs_k_fp4 = tl.arange(0, BLOCK_K // 2)
offs_bn = (pid_n * BLOCK_N + tl.arange(0, BLOCK_N)) % N
b_ptrs = b_ptr + (offs_k_fp4[:, None] * stride_bk + offs_bn[None, :] * stride_bn)
b_ptrs += k_start * (BLOCK_K // 2) * stride_bk
offs_ks = tl.arange(0, BLOCK_K // SCALE_GROUP_SIZE * 32)
offs_bsn = (pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)) % tl.cdiv(N, 32)
b_scale_ptrs = b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk
b_scale_ptrs += k_start * BLOCK_K * stride_bsk
accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for k in range(k_start, k_end):
# Load bf16 A tile
a_bf16 = tl.load(a_bf16_ptrs)
a_f32 = a_bf16.to(tl.float32)
# Use AITER's official quant function — exact same as reference
a_fp4, a_scales = _mxfp4_quant_op(a_f32, BLOCK_K, BLOCK_M, MXFP4_QUANT_BLOCK_SIZE)
# Load B (fp4, pre-quantized)
b = tl.load(b_ptrs, cache_modifier=".cg")
b_scales = tl.load(b_scale_ptrs).reshape(
BLOCK_N // 32, BLOCK_K // SCALE_GROUP_SIZE // 8, 4, 16, 2, 2, 1
).permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_N, BLOCK_K // SCALE_GROUP_SIZE)
# MXFP4 × MXFP4 GEMM via tl.dot_scaled
accumulator = tl.dot_scaled(a_fp4, a_scales, "e2m1", b, b_scales, "e2m1",
acc=accumulator, fast_math=True)
a_bf16_ptrs += BLOCK_K * stride_ak
b_ptrs += (BLOCK_K // 2) * stride_bk
b_scale_ptrs += BLOCK_K * stride_bsk
offs_cm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M).to(tl.int64)
offs_cn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N).to(tl.int64)
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
if NUM_KSPLIT == 1:
c = accumulator.to(c_ptr.type.element_ty)
c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :]
tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")
else:
c_ptrs = c_ptr + pid_k * stride_cs + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :]
tl.store(c_ptrs, accumulator, mask=c_mask)
@triton.jit
def _reduce_kernel(
partials_ptr, out_ptr, M, N,
stride_ps, stride_pm, stride_pn, stride_om, stride_on,
NUM_KSPLIT: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for s in range(NUM_KSPLIT):
ptrs = partials_ptr + s * stride_ps + offs_m[:, None] * stride_pm + offs_n[None, :] * stride_pn
acc += tl.load(ptrs, mask=mask, other=0.0)
out_ptrs = out_ptr + offs_m[:, None] * stride_om + offs_n[None, :] * stride_on
tl.store(out_ptrs, acc.to(out_ptr.type.element_ty), mask=mask)
_buffers = {}
def _get_config(m, n, k):
"""Shape-specialized configs. Fused Triton for ALL shapes."""
if m <= 16:
if k >= 4096:
return 16, 128, 256, 14, 8
elif k >= 1536:
return 16, 128, 256, 4, 8
else:
return 16, 128, 256, 1, 8
elif m <= 32:
if k >= 4096:
return 16, 64, 256, 8, 8
elif k >= 1536:
return 16, 64, 256, 4, 8
else:
return 16, 64, 256, 1, 8
elif m <= 64:
# M=64: BLOCK_M=32, 2 M-tiles. Fused quant eliminates 10µs overhead.
if k >= 4096:
return 32, 128, 256, 4, 8
elif k >= 1536:
return 32, 128, 256, 2, 8
else:
return 32, 128, 256, 1, 8
else:
# M=256: BLOCK_M=32, 8 M-tiles. Large enough for good CU fill.
if k >= 4096:
return 32, 128, 256, 2, 8
elif k >= 1536:
return 32, 128, 256, 1, 8
else:
return 32, 128, 256, 1, 8
def _triton_dispatch(A_bf16, B_q, B_scale_sh, m, n, k):
BLOCK_M, BLOCK_N, BLOCK_K, NUM_KSPLIT, num_warps = _get_config(m, n, k)
B_q_u8 = B_q.view(torch.uint8)
B_t = B_q_u8.T
B_s = B_scale_sh.view(torch.uint8)
B_scales_triton = B_s.reshape(B_s.shape[0] // 32, B_s.shape[1] * 32)
num_pid_m = triton.cdiv(m, BLOCK_M)
num_pid_n = triton.cdiv(n, BLOCK_N)
key = (m, n, k)
if key not in _buffers:
if NUM_KSPLIT == 1:
_buffers[key] = (
torch.empty((m, n), dtype=torch.bfloat16, device=A_bf16.device),
None,
)
else:
_buffers[key] = (
torch.empty((m, n), dtype=torch.bfloat16, device=A_bf16.device),
torch.empty((NUM_KSPLIT, m, n), dtype=torch.float32, device=A_bf16.device),
)
out, partials = _buffers[key]
if NUM_KSPLIT == 1:
_aiter_fused_quant_gemm_kernel[(num_pid_m * num_pid_n,)](
A_bf16, B_t, out, B_scales_triton, m, n, k,
A_bf16.stride(0), A_bf16.stride(1),
B_t.stride(0), B_t.stride(1),
0, out.stride(0), out.stride(1),
B_scales_triton.stride(0), B_scales_triton.stride(1),
BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
NUM_KSPLIT=1, MXFP4_QUANT_BLOCK_SIZE=32,
num_warps=num_warps, num_stages=2, matrix_instr_nonkdim=16,
)
else:
_aiter_fused_quant_gemm_kernel[(NUM_KSPLIT * num_pid_m * num_pid_n,)](
A_bf16, B_t, partials, B_scales_triton, m, n, k,
A_bf16.stride(0), A_bf16.stride(1),
B_t.stride(0), B_t.stride(1),
partials.stride(0), partials.stride(1), partials.stride(2),
B_scales_triton.stride(0), B_scales_triton.stride(1),
BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
NUM_KSPLIT=NUM_KSPLIT, MXFP4_QUANT_BLOCK_SIZE=32,
num_warps=num_warps, num_stages=2, matrix_instr_nonkdim=16,
)
_reduce_kernel[(triton.cdiv(m, 16), triton.cdiv(n, 64))](
partials, out, m, n,
partials.stride(0), partials.stride(1), partials.stride(2),
out.stride(0), out.stride(1),
NUM_KSPLIT=NUM_KSPLIT, BLOCK_M=16, BLOCK_N=64, num_warps=4,
)
return out
_asm_buffers = {}
MXFP4_QBS = 32
def _asm_path(A, B_shuffle, B_scale_sh):
"""v126's fast path: fused quant+shuffle + direct ASM GEMM."""
M, K = A.shape
N = B_shuffle.shape[0]
key = (M, N, K)
if key not in _asm_buffers:
scaleN_valid = triton.cdiv(K, MXFP4_QBS)
scaleN_pad = triton.cdiv(scaleN_valid, 8) * 8
sm_pad = triton.cdiv(M, 256) * 256
padded_m = (M + 31) // 32 * 32
x_fp4 = torch.empty((M, K // 2), dtype=torch.uint8, device=A.device)
scale_sh = torch.full((sm_pad * scaleN_pad,), 127, dtype=torch.uint8, device=A.device)
out = torch.empty((padded_m, N), dtype=dtypes.bf16, device=A.device)
ck_config = get_GEMM_config(M, N, K)
splitK = 0
kernelName = ""
if ck_config is not None:
splitK = ck_config.get("splitK", None)
splitK = 0 if splitK is None else splitK
kernelName = ck_config["kernelName"]
NUM_ITER, BSM, BSN, NW, NS = _get_quant_config(M, K)
_asm_buffers[key] = (x_fp4, scale_sh, scaleN_valid, scaleN_pad, sm_pad,
out, padded_m, kernelName, splitK, NUM_ITER, BSM, BSN, NW, NS)
(x_fp4, scale_sh_flat, scaleN_valid, scaleN_pad, sm_pad,
out, padded_m, kernelName, splitK, NUM_ITER, BSM, BSN, NW, NS) = _asm_buffers[key]
grid = (triton.cdiv(M, BSM), triton.cdiv(K, BSN * NUM_ITER))
_fused_quant_shuffle_kernel[grid](
A, x_fp4, scale_sh_flat, *A.stride(), *x_fp4.stride(),
M=M, N=K, scaleN_valid=scaleN_valid, scaleN_pad=scaleN_pad,
MXFP4_QUANT_BLOCK_SIZE=MXFP4_QBS, SCALING_MODE=0,
NUM_ITER=NUM_ITER, BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=BSN,
NUM_STAGES=NS, num_warps=NW, waves_per_eu=0, num_stages=1,
)
A_q = x_fp4.view(dtypes.fp4x2)
A_scale_sh = scale_sh_flat.view(sm_pad, scaleN_pad).view(dtypes.fp8_e8m0)
gemm_a4w4_asm(A_q.view(M, K // 2), B_shuffle, A_scale_sh, B_scale_sh,
out, kernelName, None, 1.0, 0.0, True, log2_k_split=splitK)
return out[:M]
def _get_quant_config(M, N):
"""Same config logic as aiter.ops.triton.quant.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))
return NUM_ITER, BLOCK_SIZE_M, BLOCK_SIZE_N, NUM_WARPS, NUM_STAGES
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]
return _triton_dispatch(A, B_q, B_scale_sh, m, n, k)
scrolls · 338 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