submission 686644
zwang86 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 278 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-686644?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:45bb2a614e64533f61c85fb1b1b24320f2a1fb6dedc1bca7574fc2ee9b18d8c0
license declaredunknown
license concludedunknown
authorszwang86
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 1
NUM_ITER, NUM_STAGES, NUM_WARPS = 1, 1, 4split-k
_SK = 15 # splitKstages = 1
NUM_WARPS, NUM_STAGES = 1, 1tile-n = 1
NUM_ITER, BLOCK_SIZE_M, BLOCK_SIZE_N = 1, triton.next_power_of_2(M), 32Kernel source
submission.py278 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
from task import input_t, output_t
import torch
import triton
import triton.language as tl
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,
gemm_a4w4_blockscale,
get_GEMM_config,
)
@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,
N,
scaleN_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_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_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
bs_n = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
d0 = bs_m // 32
d1 = (bs_m % 32) // 16
d2 = bs_m % 16
d3 = bs_n // 8
d4 = (bs_n % 8) // 4
d5 = bs_n % 4
bs_offs = (
d1[:, None]
+ d4[None, :] * 2
+ d2[:, None] * 4
+ d5[None, :] * 64
+ d3[None, :] * 256
+ d0[:, None] * 32 * scaleN_pad
)
if EVEN_M_N:
tl.store(bs_ptr + bs_offs, bs_e8m0)
else:
scaleN_valid = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
bs_mask = (bs_m[:, None] < M) & (bs_n[None, :] < scaleN_valid)
bs_val = tl.where(bs_mask, bs_e8m0, 127)
tl.store(bs_ptr + bs_offs, bs_val, mask=bs_mask)
_cfg = {}
# Config tuple indices for hot-path access
_GRID = 0
_A_FP4 = 1
_BS = 2
_SNP = 3 # scaleN_pad
_NI = 4 # NUM_ITER
_BSM = 5 # BLOCK_SIZE_M
_BSN = 6 # BLOCK_SIZE_N
_NS = 7 # NUM_STAGES
_NW = 8 # NUM_WARPS
_AQ = 9 # a_q pre-computed view
_ASC = 10 # a_sc
_OUT = 11
_OV = 12 # out_view
_UB = 13 # use_blockscale
_KN = 14 # kernelName (or None)
_SK = 15 # splitK
_K = 16
_KHALF = 17
def _init_shape(M, K, N, device):
QBLOCK = 32
a_fp4 = torch.empty((M, K // 2), dtype=torch.uint8, device=device)
scaleN_valid = triton.cdiv(K, QBLOCK)
scaleN_pad = triton.cdiv(scaleN_valid, 8) * 8
scaleM_pad = triton.cdiv(M, 256) * 256
blockscale = torch.empty(
(scaleM_pad, scaleN_pad), dtype=torch.uint8, device=device
)
a_q = a_fp4.view(dtypes.fp4x2)
a_sc = blockscale.view(dtypes.fp8_e8m0)
M_pad = (M + 31) // 32 * 32
out = torch.empty((M_pad, N), dtype=torch.bfloat16, device=device)
out_view = out[:M] if M != M_pad else out
if M <= 32:
NUM_ITER, BLOCK_SIZE_M, BLOCK_SIZE_N = 1, triton.next_power_of_2(M), 32
NUM_WARPS, NUM_STAGES = 1, 1
else:
NUM_ITER, BLOCK_SIZE_M, BLOCK_SIZE_N = 4, 64, 64
NUM_WARPS, NUM_STAGES = 4, 2
if K <= 16384:
BLOCK_SIZE_M, BLOCK_SIZE_N = 32, 128
if K <= 1024:
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 * NUM_ITER))
ck_config = get_GEMM_config(M, N, K)
use_blockscale = False
# Force 32x128 tile for all shapes (default 192x128 wastes compute for small M)
kernelName = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
# splitK for untuned shapes: improve CU utilization
tile_m, tile_n, tile_k = 32, 128, 128
tile_num = ((M_pad + tile_m - 1) // tile_m) * ((N + tile_n - 1) // tile_n)
splitK = 0
cu_num = 304 # MI355X
cus_per_tile = cu_num / max(tile_num, 1)
while (cus_per_tile >= (1 << (splitK + 1))
and (1 << (splitK + 1)) * tile_k < 2 * K
and splitK < 3):
splitK += 1
if ck_config is not None:
kn = ck_config["kernelName"]
sk = ck_config.get("splitK", None)
if "_ZN" not in kn:
use_blockscale = True
splitK = 0 if sk is None else int(sk)
else:
kernelName = kn
splitK = int(sk) if sk is not None else 0
return (
grid, # 0
a_fp4, # 1
blockscale, # 2
scaleN_pad, # 3
NUM_ITER, # 4
BLOCK_SIZE_M, # 5
BLOCK_SIZE_N, # 6
NUM_STAGES, # 7
NUM_WARPS, # 8
a_q, # 9
a_sc, # 10
out, # 11
out_view, # 12
use_blockscale,# 13
kernelName, # 14
splitK, # 15
K, # 16
K // 2, # 17
)
_introspection_cached = False
def custom_kernel(data: input_t) -> output_t:
global _introspection_cached
A = data[0]
B_sh = data[3]
B_sc = data[4]
M = A.shape[0]
K = A.shape[1]
N = B_sh.shape[0]
c = _cfg.get((M, K, N))
if c is None:
c = _init_shape(M, K, N, A.device)
_cfg[(M, K, N)] = c
_fused_quant_shuffle_kernel[c[_GRID]](
A, c[_A_FP4], c[_BS],
K, 1, c[_KHALF], 1,
M=M, N=K, scaleN_pad=c[_SNP],
MXFP4_QUANT_BLOCK_SIZE=32, SCALING_MODE=0,
NUM_ITER=c[_NI],
BLOCK_SIZE_M=c[_BSM],
BLOCK_SIZE_N=c[_BSN],
NUM_STAGES=c[_NS],
num_warps=c[_NW],
waves_per_eu=0, num_stages=1,
)
if c[_UB]:
gemm_a4w4_blockscale(
c[_AQ], B_sh, c[_ASC], B_sc, c[_OUT], splitK=c[_SK],
)
else:
gemm_a4w4_asm(
c[_AQ], B_sh, c[_ASC], B_sc, c[_OUT],
c[_KN] or "", None, 1.0, 0.0, True, c[_SK],
)
if not _introspection_cached:
_introspection_cached = True
try:
import inspect
import typing
inner = gemm_a4w4_asm.__globals__.get('_gemm_a4w4_asm')
if inner is not None and hasattr(inner, '__wrapped__'):
fn = inner.__wrapped__
fn.__signature__ = inspect.signature(fn)
_fn_id = id(fn)
_fn_hints = typing.get_type_hints(fn)
_orig_gth = typing.get_type_hints
def _fast_gth(obj, globalns=None, localns=None,
include_extras=False):
if id(obj) == _fn_id:
return _fn_hints
return _orig_gth(obj, globalns=globalns,
localns=localns,
include_extras=include_extras)
typing.get_type_hints = _fast_gth
except Exception:
pass
return c[_OV]
scrolls · 278 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