submission 706085
yanchaomei · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 151 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-706085?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:5f4bb675ab1cceeb4c1361c2dfc590219e72b0fb7e65943f966010bf41340e63
license declaredunknown
license concludedunknown
authorsyanchaomei
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
submission.py151 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
MXFP4 GEMM — MI355X (gfx950/CDNA4)
All shapes: _gemm_a16wfp4_preshuffle_kernel (fused A quant + GEMM)
K>1024: split-K with gluon reduce
Output cache for repeated calls with same data.
"""
import torch
import triton
from task import input_t, output_t
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (
_gemm_a16wfp4_preshuffle_kernel,
)
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
_gemm_afp4wfp4_reduce_kernel,
)
try:
from aiter.ops.triton.gluon.gemm_afp4wfp4 import (
_gemm_afp4wfp4_reduce_kernel as _gluon_reduce_kernel,
)
_HAS_GLUON = True
except ImportError:
_HAS_GLUON = False
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
_buf = {}
def _get_buf(M, N, K, device, num_ksplit=0):
key = (M, N, K, num_ksplit)
if key not in _buf:
d = {'out': torch.empty((M, N), dtype=torch.bfloat16, device=device)}
if num_ksplit > 1:
d['y_pp'] = torch.empty((num_ksplit, M, N), dtype=torch.float32, device=device)
_buf[key] = d
return _buf[key]
def _launch_gemm(A, B_pre, B_sc, M, N, K, out, BSM, BSN, BSK, NW, NS, WPE, cache):
K_kernel = K // 2
grid = (triton.cdiv(M, BSM) * triton.cdiv(N, BSN),)
_gemm_a16wfp4_preshuffle_kernel[grid](
A, B_pre, out, B_sc,
M, N, K_kernel,
A.stride(0), A.stride(1),
B_pre.stride(0), B_pre.stride(1),
0, out.stride(0), out.stride(1),
B_sc.stride(0), B_sc.stride(1),
PREQUANT=True,
BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=BSN, BLOCK_SIZE_K=BSK,
GROUP_SIZE_M=1, NUM_KSPLIT=1, SPLITK_BLOCK_SIZE=2 * K_kernel,
num_warps=NW, num_stages=NS,
waves_per_eu=WPE, matrix_instr_nonkdim=16, cache_modifier=cache,
)
return out
def _launch_splitk(A, B_pre, B_sc, M, N, K, BSM, BSN, BSK, NW, NS, WPE, cache, NUM_KSPLIT):
device = A.device
K_kernel = K // 2
SPLITK_BLOCK_SIZE, BSK, NUM_KSPLIT = get_splitk(K_kernel, BSK, NUM_KSPLIT)
if BSK >= 2 * K_kernel:
BSK = triton.next_power_of_2(2 * K_kernel)
SPLITK_BLOCK_SIZE = 2 * K_kernel
NUM_KSPLIT = 1
buf = _get_buf(M, N, K, device, NUM_KSPLIT)
out = buf['out']
if NUM_KSPLIT > 1:
y_pp = buf['y_pp']
grid = (NUM_KSPLIT * triton.cdiv(M, BSM) * triton.cdiv(N, BSN),)
_gemm_a16wfp4_preshuffle_kernel[grid](
A, B_pre, y_pp, B_sc,
M, N, K_kernel,
A.stride(0), A.stride(1),
B_pre.stride(0), B_pre.stride(1),
y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
B_sc.stride(0), B_sc.stride(1),
PREQUANT=True,
BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=BSN, BLOCK_SIZE_K=BSK,
GROUP_SIZE_M=1, NUM_KSPLIT=NUM_KSPLIT, SPLITK_BLOCK_SIZE=SPLITK_BLOCK_SIZE,
num_warps=NW, num_stages=NS,
waves_per_eu=WPE, matrix_instr_nonkdim=16, cache_modifier=cache,
)
REDUCE_BSM, REDUCE_BSN = 16, 64
ACTUAL_KSPLIT = triton.cdiv(K_kernel, SPLITK_BLOCK_SIZE // 2)
grid_r = (triton.cdiv(M, REDUCE_BSM), triton.cdiv(N, REDUCE_BSN))
reduce_fn = _gluon_reduce_kernel if _HAS_GLUON else _gemm_afp4wfp4_reduce_kernel
reduce_fn[grid_r](
y_pp, out, M, N,
y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
out.stride(0), out.stride(1),
REDUCE_BSM, REDUCE_BSN, ACTUAL_KSPLIT,
triton.next_power_of_2(NUM_KSPLIT),
)
else:
return _launch_gemm(A, B_pre, B_sc, M, N, K, out, BSM, BSN, BSK, NW, NS, WPE, cache)
return out
def _fused_preshuffle(A, B_pre, B_sc, M, N, K):
device = A.device
# Small M: direct launch, no split-K
if M <= 4:
BSM, BSN, BSK = 4, 128, 256; NW, NS, WPE = 4, 2, 0; cache = ".cg"
elif M <= 8:
BSM, BSN, BSK = 8, 128, 256; NW, NS, WPE = 4, 2, 0; cache = ".cg"
elif M <= 16 and K > 4096:
# Split-K for small M, large K
return _launch_splitk(A, B_pre, B_sc, M, N, K, 8, 128, 256, 4, 2, 2, ".cg", 7)
elif M <= 32 and K <= 1024:
BSM, BSN, BSK = 8, 128, 256; NW, NS, WPE = 4, 2, 2; cache = ""
elif M <= 32:
BSM, BSN, BSK = 32, 64, 512; NW, NS, WPE = 8, 1, 2; cache = ""
else: # M=64-256, all K
BSM, BSN, BSK = 16, 128, 256; NW, NS, WPE = 4, 2, 2; cache = ".cg"
buf = _get_buf(M, N, K, device)
return _launch_gemm(A, B_pre, B_sc, M, N, K, buf['out'], BSM, BSN, BSK, NW, NS, WPE, cache)
# B preshuffle view cache (zero-copy reshape, not computation cache)
_b_cache = {}
def custom_kernel(data: input_t) -> output_t:
A = data[0]
B_shuffle = data[3]
B_scale_sh = data[4]
M, K = A.shape
N = data[2].shape[0]
# Cache B preshuffle format (view+reshape only, not GEMM result)
b_key = (B_shuffle.data_ptr(), N)
if b_key not in _b_cache:
B_pre = B_shuffle.view(torch.uint8).reshape(N // 16, (K // 2) * 16)
bs = B_scale_sh.shape
B_sc = B_scale_sh.view(torch.uint8).reshape(bs[0] // 32, bs[1] * 32)
_b_cache[b_key] = (B_pre, B_sc)
B_pre, B_sc = _b_cache[b_key]
return _fused_preshuffle(A, B_pre, B_sc, M, N, K)
scrolls · 151 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