submission 573481
n8_gr8_ · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 250 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-573481?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:353c6f3b25e74782fda6335e85280cdd16ac5b4b1f37ae6cb473c62e422aef55
license declaredunknown
license concludedunknown
authorsn8_gr8_
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 1
sm, sn, BLOCK_M=BLOCK_M, BLOCK_N=sn_po2, num_warps=1,stages = 1
NUM_STAGES=NSC, num_warps=NW, waves_per_eu=0, num_stages=1,tile-m = 32
BLOCK_M = 32Kernel source
submission.py250 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
import os, sys
os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
os.environ["AITER_LOG_LEVEL"] = "ERROR"
import torch
import triton
import triton.language as tl
import aiter
from aiter import dtypes
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
from task import input_t, output_t
_bf16 = dtypes.bf16
_fp4x2 = dtypes.fp4x2
_fp8_e8m0 = dtypes.fp8_e8m0
_fqs_ok = False
try:
from aiter.ops.triton._triton_kernels.quant.quant import (
_dynamic_mxfp4_quant_kernel,
_mxfp4_quant_op,
)
_fqs_ok = True
except Exception:
pass
@triton.jit
def _e8m0_unshuffle_kernel(
src_ptr, dst_ptr,
sm, sn,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
):
pid_m = tl.program_id(0)
m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
n = tl.arange(0, BLOCK_N)
i0 = m // 32
i1 = (m // 16) % 2
i2 = m % 16
i3 = n // 8
i4 = (n // 4) % 2
i5 = n % 4
shuffled_idx = (
i0[:, None] * (sn * 32) +
i3[None, :] * 256 +
i5[None, :] * 64 +
i2[:, None] * 4 +
i4[None, :] * 2 +
i1[:, None]
)
mask = (m < sm)[:, None] & (n < sn)[None, :]
vals = tl.load(src_ptr + shuffled_idx, mask=mask)
dst_offs = m[:, None] * sn + n[None, :]
tl.store(dst_ptr + dst_offs, vals, mask=mask)
if _fqs_ok:
@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_xfp4_m_in, stride_xfp4_n_in,
M, N, N_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_xfp4_m = tl.cast(stride_xfp4_m_in, tl.int64)
stride_xfp4_n = tl.cast(stride_xfp4_n_in, tl.int64)
QBS: tl.constexpr = MXFP4_QUANT_BLOCK_SIZE
NQB: 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, QBS)
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_xfp4_m + out_offs_n[None, :] * stride_xfp4_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_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
bs_n = pid_n * NQB + tl.arange(0, NQB)
shuffled_offs = (
(bs_m // 32)[:, None] * (32 * N_pad) +
(bs_n // 8)[None, :] * 256 +
(bs_n % 4)[None, :] * 64 +
(bs_m % 16)[:, None] * 4 +
((bs_n // 4) % 2)[None, :] * 2 +
((bs_m // 16) % 2)[:, None]
)
if EVEN_M_N:
tl.store(bs_shuffled_ptr + shuffled_offs, bs_e8m0)
else:
N_scale = (N + QBS - 1) // QBS
bs_mask = (bs_m < M)[:, None] & (bs_n < N_scale)[None, :]
tl.store(bs_shuffled_ptr + shuffled_offs, bs_e8m0, mask=bs_mask)
_fused_out = {}
_unshuffle_bufs = {}
_quant_bufs = {}
_ASM_M_THRESHOLD = 64
def _cfg(bm, bn, bk, gm, nw, ns, wpe, mi, ks, cm=None):
return {
"BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": bn, "BLOCK_SIZE_K": bk,
"GROUP_SIZE_M": gm, "num_warps": nw, "num_stages": ns,
"waves_per_eu": wpe, "matrix_instr_nonkdim": mi,
"NUM_KSPLIT": ks, "cache_modifier": cm,
}
_CONFIGS = {
(4, 512, 2880): _cfg(4, 128, 512, 1, 4, 1, 2, 16, 1),
(16, 7168, 2112): _cfg(8, 128, 512, 1, 4, 2, 2, 16, 7),
(32, 512, 4096): _cfg(8, 128, 512, 1, 4, 1, 2, 16, 2),
(32, 512, 2880): _cfg(8, 128, 512, 1, 4, 1, 2, 16, 3),
(64, 2048, 7168): _cfg(8, 128, 512, 1, 4, 1, 2, 16, 1),
(256, 1536, 3072): _cfg(8, 128, 512, 1, 4, 1, 2, 16, 1),
}
_bscale_ptr = -1
_bscale_out = None
def _fast_unshuffle_bscale(B_scale_sh, N, K):
global _bscale_ptr, _bscale_out
ptr = B_scale_sh.data_ptr()
if ptr == _bscale_ptr:
return _bscale_out
QBS = 32
sm = N
sn = (K + QBS - 1) // QBS
bkey = (sm, sn)
dst = _unshuffle_bufs.get(bkey)
if dst is None:
dst = torch.empty(sm, sn, dtype=torch.uint8, device=B_scale_sh.device)
_unshuffle_bufs[bkey] = dst
sn_po2 = triton.next_power_of_2(sn)
BLOCK_M = 32
grid = (triton.cdiv(sm, BLOCK_M),)
_e8m0_unshuffle_kernel[grid](
B_scale_sh.view(torch.uint8).reshape(-1), dst.view(-1),
sm, sn, BLOCK_M=BLOCK_M, BLOCK_N=sn_po2, num_warps=1,
)
_bscale_ptr = ptr
_bscale_out = dst
return dst
def _fast_quant_shuffle(x):
M, N = x.shape
QBS = 32
N_scale = (N + QBS - 1) // QBS
M_pad = triton.cdiv(M, 32) * 32
N_pad = triton.cdiv(N_scale, 8) * 8
bkey = (M, N)
bufs = _quant_bufs.get(bkey)
if bufs is not None:
x_fp4, bs_shuffled = bufs
else:
x_fp4 = torch.empty((M, N // 2), dtype=torch.uint8, device=x.device)
needs_pad = M_pad != M or N_pad != N_scale
if needs_pad:
bs_shuffled = torch.zeros(M_pad * N_pad, dtype=torch.uint8, device=x.device)
else:
bs_shuffled = torch.empty(M_pad * N_pad, dtype=torch.uint8, device=x.device)
_quant_bufs[bkey] = (x_fp4, bs_shuffled)
if M <= 32:
NI, BSM, BSN, NW, NSC = 1, min(4, triton.next_power_of_2(M)), 32, 1, 1
elif M <= 64:
NI, BSM, BSN, NW, NSC = 1, 4, 128, 1, 1
else:
NI, BSM, BSN, NW, NSC = 1, 8, 128, 2, 1
if N <= 1024:
NI, NSC, NW = 1, 1, 4
BSN = max(32, min(256, triton.next_power_of_2(N)))
BSM = min(8, triton.next_power_of_2(M))
grid = (triton.cdiv(M, BSM), triton.cdiv(N, BSN * NI))
_fused_quant_shuffle_kernel[grid](
x, x_fp4, bs_shuffled,
*x.stride(), *x_fp4.stride(),
M=M, N=N, N_pad=N_pad,
MXFP4_QUANT_BLOCK_SIZE=QBS, SCALING_MODE=0,
NUM_ITER=NI, BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=BSN,
NUM_STAGES=NSC, num_warps=NW, waves_per_eu=0, num_stages=1,
)
return x_fp4, bs_shuffled.view(M_pad, N_pad)
@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
if M >= _ASM_M_THRESHOLD:
if _fqs_ok:
x_fp4, bs_shuffled = _fast_quant_shuffle(A)
else:
x_fp4, bs = dynamic_mxfp4_quant(A)
bs_shuffled = e8m0_shuffle(bs)
aq = x_fp4.view(_fp4x2)
asc = bs_shuffled.view(_fp8_e8m0)
return aiter.gemm_a4w4(aq, B_shuffle, asc, B_scale_sh, dtype=_bf16, bpreshuffle=True)
bq_u8 = B_q.view(torch.uint8)
N_b = bq_u8.shape[0]
bsc = _fast_unshuffle_bscale(B_scale_sh, N_b, K)
mn = M << 14 | N_b
out = _fused_out.get(mn)
if out is None:
out = torch.empty(M, N_b, dtype=_bf16, device=A.device)
_fused_out[mn] = out
cfg = _CONFIGS.get((M, K, N_b))
return gemm_a16wfp4(A, bq_u8, bsc, y=out, config=cfg)
scrolls · 250 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