submission 693087
LiangSu8899 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 206 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-693087?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:53859815c26e0300b9cea659ffb40159329b5a7ae3cab7e9e27ec66e6ae492bd
license declaredunknown
license concludedunknown
authorsLiangSu8899
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
split-k
custom_kernel._get_splitk = gmod.get_splitkstages = 1
num_warps=NW_q, waves_per_eu=0, num_stages=1)Kernel source
submission.py206 lines
"""GEMM v162: Manually tuned fused kernel configs for k<=512 shapes.
v158 diagnostic showed default _get_config returns:
m=4: BSM=4, BSN=128, warps=4, grid=23
m=32: BSM=8, BSN=128, warps=8, grid=128/92
m=64: BSM=16, BSN=128, warps=8, stages=2, grid=224
m=256: BSM=8, BSN=128, warps=8, grid=768
Try BSN=64 for more parallelism, fewer warps for less overhead.
Keep quant+ASM for k>512 (proven best).
"""
from task import input_t, output_t
import triton
def custom_kernel(data: input_t) -> output_t:
import torch
import aiter
from aiter import dtypes
from aiter.ops.triton.gemm.basic import gemm_a16wfp4 as gmod
A, B, B_q, B_shuffle, B_scale_sh = data
m, k = A.shape
n = B.shape[0]
if not hasattr(custom_kernel, '_init'):
custom_kernel._init = True
custom_kernel._b = {}
custom_kernel._fused_kernel = gmod._gemm_a16wfp4_preshuffle_kernel
custom_kernel._get_splitk = gmod.get_splitk
custom_kernel._get_config = gmod._get_config
custom_kernel._fp4x2 = dtypes.fp4x2
custom_kernel._e8m0 = dtypes.fp8_e8m0
from aiter.ops.triton.quant.quant import _mxfp4_quant_op
import triton.language as tl
@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 _fq(
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, SN,
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_x_fp4_m_in, tl.int64)
stride_fp4_n = tl.cast(stride_x_fp4_n_in, tl.int64)
NUM_QB: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
for pid_n in tl.range(start_n, tl.minimum(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_t, bs_e8 = _mxfp4_quant_op(x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE)
o_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
o_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
o_offs = o_m[:, None] * stride_fp4_m + o_n[None, :] * stride_fp4_n
if EVEN_M_N:
tl.store(x_fp4_ptr + o_offs, out_t)
else:
tl.store(x_fp4_ptr + o_offs, out_t, mask=(o_m < M)[:, None] & (o_n < N // 2)[None, :])
bs_r = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
bs_c = pid_n * NUM_QB + tl.arange(0, NUM_QB)
n_sc = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
d0 = bs_r // 32; d5 = (bs_r % 32) // 16; d3 = bs_r % 16
d1 = bs_c // 8; d4 = (bs_c % 8) // 4; d2 = bs_c % 4
SN64 = tl.cast(SN, tl.int64)
sh_offs = (d0[:, None].to(tl.int64) * (SN64 * 32)
+ d1[None, :].to(tl.int64) * 256
+ d2[None, :].to(tl.int64) * 64
+ d3[:, None].to(tl.int64) * 4
+ d4[None, :].to(tl.int64) * 2
+ d5[:, None].to(tl.int64))
if EVEN_M_N:
tl.store(bs_ptr + sh_offs, bs_e8)
else:
tl.store(bs_ptr + sh_offs, bs_e8, mask=(bs_r < M)[:, None] & (bs_c < n_sc)[None, :])
custom_kernel._fq = _fq
custom_kernel._asm = aiter.gemm_a4w4_asm
b = custom_kernel._b
key = (m, n, k)
if key not in b:
if k <= 512:
K_val = k // 2
# Manually tuned configs per shape
# Default from _get_config: BSM=4/8, BSN=128, BSK=512
# Try BSN=64 for more tiles, lower warps for less overhead
if m <= 4:
BSM, BSN, BSK = 4, 64, 512
GSM, NKS = 1, 1
nw, ns, wpe, cm = 4, 1, 2, '.cg'
elif m <= 16:
BSM, BSN, BSK = 4, 128, 512
GSM, NKS = 1, 1
nw, ns, wpe, cm = 4, 1, 2, '.cg'
elif m <= 32:
BSM, BSN, BSK = 8, 64, 512
GSM, NKS = 1, 1
nw, ns, wpe, cm = 4, 1, 2, '.cg'
else:
BSM, BSN, BSK = 8, 128, 512
GSM, NKS = 1, 1
nw, ns, wpe, cm = 8, 1, 2, '.cg'
SPLITK_BS, BSK, NKS = custom_kernel._get_splitk(K_val, BSK, NKS)
grid_mn = triton.cdiv(m, BSM) * triton.cdiv(n, BSN)
out = torch.empty(m, n, dtype=torch.bfloat16, device=A.device)
EVEN_K = (K_val % (BSK // 2) == 0 and SPLITK_BS % BSK == 0 and K_val % (SPLITK_BS // 2) == 0)
meta = {
'BLOCK_SIZE_M': BSM, 'BLOCK_SIZE_N': BSN, 'BLOCK_SIZE_K': BSK,
'GROUP_SIZE_M': GSM, 'NUM_KSPLIT': NKS, 'SPLITK_BLOCK_SIZE': SPLITK_BS,
'num_warps': nw, 'num_stages': ns, 'waves_per_eu': wpe,
'matrix_instr_nonkdim': 16, 'cache_modifier': cm,
'GRID_MN': grid_mn, 'PREQUANT': True,
}
b[key] = ('fused', out, K_val, EVEN_K, meta, grid_mn, NKS)
else:
SG = 32
n_sc = (k + SG - 1) // SG
sm = (m + 255) // 256 * 256
sn = (n_sc + 7) // 8 * 8
BSM_q = 16
BSN_q = 32
NI_q, NW_q, NS_q = 1, 2, 2
if m >= 64:
BSM_q = 64
grid_q = (triton.cdiv(m, BSM_q), triton.cdiv(k, BSN_q * NI_q))
x_fp4 = torch.empty((m, k // 2), dtype=torch.uint8, device=A.device)
scale_sh = torch.zeros(sm * sn, dtype=torch.uint8, device=A.device)
A_q_view = x_fp4.view(custom_kernel._fp4x2)
A_s_view = scale_sh.view(sm, sn).view(custom_kernel._e8m0)
out = torch.empty(m, n, dtype=torch.bfloat16, device=A.device)
tile_m, tile_n = 32, 128
base_name = f"f4gemm_bf16_per1x32Fp4_BpreShuffle_{tile_m}x{tile_n}"
kname = f"_ZN5aiter{len(base_name)}{base_name}E"
b[key] = ('asm', x_fp4, scale_sh, grid_q,
BSM_q, BSN_q, NI_q, NS_q, NW_q,
sm, sn, A_q_view, A_s_view, out, kname)
entry = b[key]
if entry[0] == 'fused':
_, out, K_val, EVEN_K, meta, grid_mn, NKS = entry
kernel = custom_kernel._fused_kernel
b_data = B_shuffle.view(torch.uint8)
b_scale = B_scale_sh.view(torch.uint8)
stride_bn = 16 * b_data.stride(0)
stride_bk = b_data.stride(1)
stride_bsn = 32 * b_scale.stride(0)
stride_bsk = b_scale.stride(1)
stride_ck, stride_cm, stride_cn = 0, out.stride(0), out.stride(1)
grid = (grid_mn * NKS,)
kernel[grid](
A, b_data, out, b_scale,
m, n, K_val,
A.stride(0), A.stride(1),
stride_bn, stride_bk,
stride_ck, stride_cm, stride_cn,
stride_bsn, stride_bsk,
EVEN_K=EVEN_K, **meta,
)
if NKS > 1:
return out.sum(dim=0).to(torch.bfloat16)
return out
else:
(_, x_fp4, scale_sh, grid_q,
BSM_q, BSN_q, NI_q, NS_q, NW_q,
sm, sn, A_q_view, A_s_view, out, kname) = entry
custom_kernel._fq[grid_q](
A, x_fp4, scale_sh,
A.stride(0), A.stride(1),
x_fp4.stride(0), x_fp4.stride(1),
M=m, N=k, SN=sn,
BLOCK_SIZE_M=BSM_q, BLOCK_SIZE_N=BSN_q,
NUM_ITER=NI_q, NUM_STAGES=NS_q,
MXFP4_QUANT_BLOCK_SIZE=32, SCALING_MODE=0,
num_warps=NW_q, waves_per_eu=0, num_stages=1)
custom_kernel._asm(
A_q_view, B_shuffle, A_s_view, B_scale_sh,
out, kname, bpreshuffle=True)
return out
scrolls · 206 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