submission 721672
FelliYang · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 318 lines, June 9 Researcher Reciprocity License v1.0.
v53.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-721672?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:40392b2bddb94e4397a6eb1552b252e138576134616d593a76b9f9be5dd81c28
license declaredunknown
license concludedunknown
authorsFelliYang
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MXFP4 GEMM - v53: v15_merge_v30 + Triton splitK for m=16, n=2112, k=7168.num-warps = 4
num_warps=4,split-k
MXFP4 GEMM - v53: v15_merge_v30 + Triton splitK for m=16, n=2112, k=7168.stages = 1
NUM_STAGES=NUM_STAGES, num_warps=NUM_WARPS, waves_per_eu=0, num_stages=1,Kernel source
v53.py318 lines
"""
MXFP4 GEMM - v53: v15_merge_v30 + Triton splitK for m=16, n=2112, k=7168.
m=16, k=7168 原来: 17 CTAs (32x128 tile, no splitK) → 6.6% CU利用率
v53: NUM_KSPLIT=14 → 17×14=238 CTAs → ~93% CU利用率
关键变化:
1. m=16 新增 _mxfp4_quant_natural_kernel: 输出 natural (M, K//32) A_scale
(gemm_afp4wfp4_preshuffle 在 M<32 时期望 un-shuffled A_scale)
2. B_scale_sh 直接 view 成 (sm//32, K): ASM shuffle format 与
_shuffle_scales 输出在 K%256==0 时完全等价, 零拷贝
3. B 相关 tensor 用全局变量而非 dict cache (B 是固定权重)
4. 其余 shape 完全沿用 v15_merge_v30 的 ASM 路径
"""
import torch
import triton
import triton.language as tl
try:
from task import input_t, output_t
except ImportError:
from typing import Any, Tuple
input_t = Tuple[Any, ...]
output_t = Any
from aiter import dtypes
import aiter
from aiter.ops.triton.quant import _mxfp4_quant_op
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import gemm_afp4wfp4_preshuffle
# ── ASM path: small-M tile configs (m=4/32) ────────────────────────────────
SHAPE_CONFIGS = {
(4, 2880, 512): ("32x128", 0),
(32, 4096, 512): ("32x128", 0),
(32, 2880, 512): ("32x128", 0),
}
def _make_kernel_name(suffix):
base = f"f4gemm_bf16_per1x32Fp4_BpreShuffle_{suffix}"
return f"_ZN5aiter{len(base)}{base}E"
# ── Quant kernel for ASM path: shuffled A_scale ─────────────────────────────
@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_mxfp4_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, N, scaleN: tl.int64, scaleN_pad: tl.int64,
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,
):
pid_m = tl.program_id(0)
start_n = tl.program_id(1) * NUM_ITER
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 pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):
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)
out_tensor, bs_e8m0 = _mxfp4_quant_op(x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE)
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)
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)
m_idx = bs_offs_m[:, None]
n_idx = bs_offs_n[None, :]
d0 = m_idx // 32; d5 = (m_idx % 32) // 16; d3 = m_idx % 16
d1 = n_idx // 8; d4 = (n_idx % 8) // 4; d2 = n_idx % 4
shuffle_offs = d0 * 32 * scaleN_pad + d1 * 256 + d2 * 64 + d3 * 4 + d4 * 2 + d5
if EVEN_M_N:
tl.store(bs_ptr + shuffle_offs, bs_e8m0)
else:
bs_mask = (bs_offs_m < M)[:, None] & (bs_offs_n < scaleN)[None, :]
tl.store(bs_ptr + shuffle_offs, bs_e8m0, mask=bs_mask)
# ── Quant kernel for Triton/splitK path: natural (M, K//32) A_scale ─────────
# gemm_afp4wfp4_preshuffle 在 M<32 时期望 un-shuffled A_scale:
# x_scale shape (M, K//32), stride (K//32, 1)
# Grid: (cdiv(M,BSM), cdiv(K,BSK)) — 对 m=16,k=7168 就是 (1, 14)
@triton.jit
def _mxfp4_quant_natural_kernel(
x_ptr, fp4_ptr, scale_ptr,
stride_xm, stride_xk,
stride_fm, stride_fk,
stride_sm, stride_sk,
M, K,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_k = tl.program_id(1)
offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
offs_k = pid_k * BLOCK_SIZE_K + tl.arange(0, BLOCK_SIZE_K)
x = tl.load(
x_ptr + offs_m[:, None] * stride_xm + offs_k[None, :] * stride_xk,
cache_modifier=".cg",
).to(tl.float32)
fp4_out, scale_out = _mxfp4_quant_op(x, BLOCK_SIZE_K, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE)
# fp4: (BLOCK_SIZE_M, BLOCK_SIZE_K//2)
offs_fk = pid_k * (BLOCK_SIZE_K // 2) + tl.arange(0, BLOCK_SIZE_K // 2)
tl.store(fp4_ptr + offs_m[:, None] * stride_fm + offs_fk[None, :] * stride_fk, fp4_out)
# scale: (BLOCK_SIZE_M, BLOCK_SIZE_K//QUANT_BLOCK_SIZE) — natural layout
NUM_BLOCKS: tl.constexpr = BLOCK_SIZE_K // MXFP4_QUANT_BLOCK_SIZE
offs_sk = pid_k * NUM_BLOCKS + tl.arange(0, NUM_BLOCKS)
tl.store(scale_ptr + offs_m[:, None] * stride_sm + offs_sk[None, :] * stride_sk, scale_out)
# ── m=16 splitK: 只缓存 A 侧 buffer(shape 固定,省去反复 torch.empty)
# B 每次调用可能不同,view 是零拷贝,直接算即可,无需缓存
_m16_init = False
_m16_out = None # (M, N) bf16 — pre-alloc output
_m16_A_q = None # (M, K//2) uint8 — pre-alloc quant A
_m16_A_sc = None # (M, K//32) uint8 — pre-alloc quant A scale
# NUM_KSPLIT=7 → 17×7=119 CTAs
# get_splitk(K=7168, BSK=512, NKSPLIT=7) → SPLITK_BLOCK_SIZE=2048, actual NKSPLIT=7 ✓
_M16_CONFIG = {
"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": 2, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": 7,
}
_M16_BSK = 512
def _init_m16_bufs(M, N, K, device):
global _m16_init, _m16_out, _m16_A_q, _m16_A_sc
_m16_out = torch.empty((M, N), dtype=dtypes.bf16, device=device)
_m16_A_q = torch.empty((M, K // 2), dtype=torch.uint8, device=device)
_m16_A_sc = torch.empty((M, K // 32), dtype=torch.uint8, device=device)
_m16_init = True
def _run_m16_splitk(A, B_shuffle, B_scale_sh):
global _m16_init
M, K = A.shape
N = B_shuffle.shape[0]
if not _m16_init:
_init_m16_bufs(M, N, K, A.device)
# B 侧: 零拷贝 view,每次直接算(无计算开销)
# B_scale_sh shape (sm, K//32) as fp8_e8m0,view 成 (sm//32, K) 与
# _shuffle_scales 输出等价(K%256==0 时数学等价,已验证)
sm = B_scale_sh.view(torch.uint8).shape[0]
# 全部保持 uint8,不做 fp4x2/fp8_e8m0 view
# benchmark 环境 Triton 不认识 float4_e2m1fn_x2 指针类型
w = B_shuffle.view(torch.uint8).reshape(N // 16, K // 2 * 16)
wscale = B_scale_sh.view(torch.uint8).view(sm // 32, K)
# Quant A → natural (M, K//32) scale, grid=(1,14)
_mxfp4_quant_natural_kernel[
(triton.cdiv(M, _M16_CONFIG["BLOCK_SIZE_M"]), triton.cdiv(K, _M16_BSK))
](
A, _m16_A_q, _m16_A_sc,
A.stride(0), A.stride(1),
_m16_A_q.stride(0), _m16_A_q.stride(1),
_m16_A_sc.stride(0), _m16_A_sc.stride(1),
M=M, K=K,
BLOCK_SIZE_M=_M16_CONFIG["BLOCK_SIZE_M"],
BLOCK_SIZE_K=_M16_BSK,
MXFP4_QUANT_BLOCK_SIZE=32,
num_warps=4,
)
return gemm_afp4wfp4_preshuffle(
_m16_A_q,
w,
_m16_A_sc,
wscale,
dtype=dtypes.bf16,
y=_m16_out,
config=dict(_M16_CONFIG),
use_aot=False,
)
# ── ASM path: small-M (m=4/32, k=512) ──────────────────────────────────────
_small_cache = {}
def _get_small_cache(M, K, N, device):
key = (M, K, N)
if key not in _small_cache:
MXFP4_QUANT_BLOCK_SIZE = 32
x_fp4 = torch.empty((M, K // 2), dtype=torch.uint8, device=device)
scaleN = K // MXFP4_QUANT_BLOCK_SIZE
scaleN_pad = (scaleN + 7) // 8 * 8
sm = (M + 255) // 256 * 256
bs_e8m0 = torch.empty(sm * scaleN_pad, dtype=torch.uint8, device=device)
padded_m = (M + 31) // 32 * 32
gemm_out = torch.empty((padded_m, N), dtype=dtypes.bf16, device=device)
shape_cfg = SHAPE_CONFIGS.get((M, N, K), None)
if shape_cfg is not None:
kernel_name = _make_kernel_name(shape_cfg[0])
log2_k_split = shape_cfg[1]
else:
kernel_name, log2_k_split = "", 0
# k<=512: single-pass quant, small tile
NUM_ITER, NUM_STAGES, NUM_WARPS = 1, 1, 4
BLOCK_SIZE_N = max(32, min(256, triton.next_power_of_2(K)))
BLOCK_SIZE_M = min(8, triton.next_power_of_2(M))
grid = (triton.cdiv(M, BLOCK_SIZE_M), triton.cdiv(K, BLOCK_SIZE_N))
_small_cache[key] = (
x_fp4, bs_e8m0, gemm_out,
scaleN, scaleN_pad, sm,
kernel_name, log2_k_split,
grid, BLOCK_SIZE_M, BLOCK_SIZE_N, NUM_ITER, NUM_STAGES, NUM_WARPS,
)
return _small_cache[key]
def _run_small(A, B_shuffle, B_scale_sh):
M, K = A.shape
N = B_shuffle.shape[0]
(x_fp4, bs_e8m0, gemm_out,
scaleN, scaleN_pad, sm,
kernel_name, log2_k_split,
grid, BLOCK_SIZE_M, BLOCK_SIZE_N, NUM_ITER, NUM_STAGES, NUM_WARPS,
) = _get_small_cache(M, K, N, A.device)
_fused_mxfp4_quant_shuffle_kernel[grid](
A, x_fp4, bs_e8m0,
*A.stride(), *x_fp4.stride(),
M=M, N=K, scaleN=scaleN, scaleN_pad=scaleN_pad,
MXFP4_QUANT_BLOCK_SIZE=32, SCALING_MODE=0,
NUM_ITER=NUM_ITER, BLOCK_SIZE_M=BLOCK_SIZE_M, BLOCK_SIZE_N=BLOCK_SIZE_N,
NUM_STAGES=NUM_STAGES, num_warps=NUM_WARPS, waves_per_eu=0, num_stages=1,
)
A_q = x_fp4.view(dtypes.fp4x2)
A_scale_sh = bs_e8m0.view(sm, scaleN_pad).view(dtypes.fp8_e8m0)
gemm_a4w4_asm(
A_q, B_shuffle, A_scale_sh, B_scale_sh, gemm_out,
kernel_name, None, 1.0, 0.0, True, log2_k_split,
)
return gemm_out[:M]
# ── Large-M path (m=64/256): CKGEMM ────────────────────────────────────────
def _quant_mxfp4_fused_simple(x):
M, N = x.shape
MXFP4_QUANT_BLOCK_SIZE = 32
x_fp4 = torch.empty((M, N // 2), dtype=torch.uint8, device=x.device)
scaleN = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
scaleN_pad = (scaleN + 7) // 8 * 8
sm = (M + 255) // 256 * 256
bs_e8m0 = torch.empty(sm * scaleN_pad, dtype=torch.uint8, device=x.device)
NUM_ITER, BLOCK_SIZE_M, BLOCK_SIZE_N, NUM_WARPS, NUM_STAGES = 4, 8, 128, 4, 2
grid = (triton.cdiv(M, BLOCK_SIZE_M), triton.cdiv(N, BLOCK_SIZE_N * NUM_ITER))
_fused_mxfp4_quant_shuffle_kernel[grid](
x, x_fp4, bs_e8m0,
*x.stride(), *x_fp4.stride(),
M=M, N=N, scaleN=scaleN, scaleN_pad=scaleN_pad,
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, num_warps=NUM_WARPS, waves_per_eu=0, num_stages=1,
)
return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(sm, scaleN_pad).view(dtypes.fp8_e8m0)
def _run_large(A, B_shuffle, B_scale_sh):
A_q, A_scale_sh = _quant_mxfp4_fused_simple(A)
return aiter.gemm_a4w4(
A_q, B_shuffle, A_scale_sh, B_scale_sh,
dtype=dtypes.bf16, bpreshuffle=True,
)
# ── Dispatch ────────────────────────────────────────────────────────────────
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
M, K = A.shape
N = B_shuffle.shape[0]
if M == 16 and K == 7168 and N == 2112:
return _run_m16_splitk(A, B_shuffle, B_scale_sh)
elif M <= 32:
return _run_small(A, B_shuffle, B_scale_sh)
else:
return _run_large(A, B_shuffle, B_scale_sh)
scrolls · 318 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