submission 645164
Esquie · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 161 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-645164?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:ee6e042eff0b8bf08335930717e3071ab7949708201ed7f311244b997361d20b
license declaredunknown
license concludedunknown
authorsEsquie
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MXFP4 GEMM: fused quant+shuffle + explicit ASM kernel selection.split-k
kernel_name, split_k = GEMM_CONFIGS.get((m, n, k), (ASM_32x128, 0))stages = 1
num_warps=c["NW"], waves_per_eu=0, num_stages=1,Kernel source
submission.py161 lines
"""
MXFP4 GEMM: fused quant+shuffle + explicit ASM kernel selection.
Default path picks 192x128 for small M (wastes 188 rows at M=4).
Force 32x128 for all shapes — matches tuned config from 256-CU CSV.
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
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
_cache = {}
SCALE_GROUP_SIZE = 32
ASM_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
ASM_64x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E"
@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_shuffled_ptr,
stride_x_m_in, stride_x_n_in,
stride_fp4_m_in, stride_fp4_n_in,
M, N, K_SCALE_PAD,
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_fp4_m = tl.cast(stride_fp4_m_in, tl.int64)
stride_fp4_n = tl.cast(stride_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_fp4_m + out_offs_n[None, :] * stride_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)
rows = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)).to(tl.int64)
cols = (pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)).to(tl.int64)
i0 = rows // 32; i1 = (rows // 16) % 2; i2 = rows % 16
i3 = cols // 8; i4 = (cols // 4) % 2; i5 = cols % 4
K_SP = tl.cast(K_SCALE_PAD, tl.int64)
shuffled_idx = (i0[:, None] * (K_SP * 32) + i3[None, :] * 256
+ i5[None, :] * 64 + i2[:, None] * 4
+ i4[None, :] * 2 + i1[:, None])
if EVEN_M_N:
tl.store(bs_shuffled_ptr + shuffled_idx, bs_e8m0)
else:
n_scale = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
bs_mask = (rows[:, None] < M) & (cols[None, :] < n_scale)
tl.store(bs_shuffled_ptr + shuffled_idx, bs_e8m0, mask=bs_mask)
def _get_quant_params(m, k):
if m <= 32:
NI, BSM, BSN, NW, NS = 1, triton.next_power_of_2(m), 32, 1, 1
else:
NI, BSM, BSN, NW, NS = 4, 64, 64, 4, 2
if k <= 16384:
BSM, BSN = 32, 128
if k <= 1024:
NI, NS, NW = 1, 1, 4
BSN = max(32, min(256, triton.next_power_of_2(k)))
BSM = min(8, triton.next_power_of_2(m))
grid = (triton.cdiv(m, BSM), triton.cdiv(k, BSN * NI))
return grid, BSM, BSN, NI, NS, NW
# Per-shape GEMM config: (kernel_name, log2_k_split)
# 32x128 is the best ASM kernel for M<=64 per the 256-CU tuned CSV
# 64x128 is better for M=256 per the CSV
GEMM_CONFIGS = {
(4, 2880, 512): (ASM_32x128, 0),
(16, 2112, 7168): (ASM_32x128, 0),
(32, 4096, 512): (ASM_32x128, 0),
(32, 2880, 512): (ASM_32x128, 0),
(64, 7168, 2048): (ASM_32x128, 0),
(256, 3072, 1536): (ASM_32x128, 0),
}
def _setup(m, n, k):
k_scale = k // SCALE_GROUP_SIZE
m_pad = ((m + 255) // 256) * 256
k_scale_pad = ((k_scale + 7) // 8) * 8
x_fp4 = torch.empty((m, k // 2), dtype=torch.uint8, device="cuda")
bs_shuffled = torch.empty(m_pad * k_scale_pad, dtype=torch.uint8, device="cuda")
out = torch.empty(((m + 31) // 32 * 32, n), dtype=torch.bfloat16, device="cuda")
grid, BSM, BSN, NI, NS, NW = _get_quant_params(m, k)
kernel_name, split_k = GEMM_CONFIGS.get((m, n, k), (ASM_32x128, 0))
return {
"x_fp4": x_fp4, "bs_shuffled": bs_shuffled, "out": out,
"m_pad": m_pad, "k_scale_pad": k_scale_pad,
"grid": grid, "BSM": BSM, "BSN": BSN, "NI": NI, "NS": NS, "NW": NW,
"kernel_name": kernel_name, "split_k": split_k,
}
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
m, k = A.shape
n = B.shape[0]
key = (m, n, k)
if key not in _cache:
_cache[key] = _setup(m, n, k)
c = _cache[key]
_fused_quant_shuffle_kernel[c["grid"]](
A, c["x_fp4"], c["bs_shuffled"],
A.stride(0), A.stride(1),
c["x_fp4"].stride(0), c["x_fp4"].stride(1),
m, k, c["k_scale_pad"],
BLOCK_SIZE_M=c["BSM"], BLOCK_SIZE_N=c["BSN"],
NUM_ITER=c["NI"], NUM_STAGES=c["NS"],
MXFP4_QUANT_BLOCK_SIZE=32, SCALING_MODE=0,
num_warps=c["NW"], waves_per_eu=0, num_stages=1,
)
A_q = c["x_fp4"].view(dtypes.fp4x2)
A_scale_sh = c["bs_shuffled"].view(c["m_pad"], c["k_scale_pad"]).view(dtypes.fp8_e8m0)
gemm_a4w4_asm(
A_q, B_shuffle, A_scale_sh, B_scale_sh,
c["out"], c["kernel_name"],
None, 1.0, 0.0, True, c["split_k"],
)
return c["out"][:m]scrolls · 161 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