submission 531060
josusanmartin · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 153 lines, June 9 Researcher Reciprocity License v1.0.
mxfp4_v103_best_combined.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-531060?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:4d3088fe3c3cd08224c01efe534078fe44a29d2ec8755d4b8d990be05ee6f03c
license declaredunknown
license concludedunknown
authorsjosusanmartin
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
split-k
- M=64: ASM splitK=2 → 13.7µsstages = 2
- M=4: BLOCK_SIZE_N=64 (exact fit for N=2880), num_stages=2 → 6.64µs (was 7.24)tile-n = 64
Version 103: Best combined - v99 + v101's M=4 BLOCK_SIZE_N=64 improvement.Kernel source
mxfp4_v103_best_combined.py153 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
Version 103: Best combined - v99 + v101's M=4 BLOCK_SIZE_N=64 improvement.
- M=4: BLOCK_SIZE_N=64 (exact fit for N=2880), num_stages=2 → 6.64µs (was 7.24)
- M=16: BLOCK_SIZE_N=64, num_stages=2, KSPLIT=7 → 14.2µs
- M=32: Default fused (no explicit config) → 9.5µs
- M=64: ASM splitK=2 → 13.7µs
- M=256: ASM splitK=1 → 12.5µs
Expected geomean: ~10.5µs
"""
import torch
import triton
import triton.language as tl
import aiter
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
from aiter.utility.fp4_utils import _dynamic_mxfp4_quant_kernel_asm_layout
from task import input_t, output_t
@triton.jit
def _mxfp4_quant_op_asm_exact(
x,
BLOCK_SIZE_N,
BLOCK_SIZE_M,
MXFP4_QUANT_BLOCK_SIZE,
):
E8_BIAS: tl.constexpr = 127
E2_BIAS: tl.constexpr = 1
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
amax = amax.to(tl.int32, bitcast=True)
amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax = amax.to(tl.float32, bitcast=True)
scale_e8m0_unbiased = tl.log2(amax).floor() - 2
scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
quant_scale = tl.exp2(-scale_e8m0_unbiased)
qx = x * quant_scale
qx = qx.to(tl.uint32, bitcast=True)
s = qx & 0x80000000
e = (qx >> 23) & 0xFF
m = qx & 0x7FFFFF
adjusted_exponents = tl.core.sub(E8_BIAS, e + 1, sanitize_overflow=False)
m = tl.where(e < E8_BIAS, (0x400000 | (m >> 1)) >> adjusted_exponents, m)
e = tl.maximum(e, E8_BIAS - E2_BIAS) - (E8_BIAS - E2_BIAS)
e2m1_tmp = tl.minimum((((e << 2) | (m >> 21)) + 1) >> 1, 0x7)
e2m1_value = ((s >> 28) | e2m1_tmp).to(tl.uint8)
e2m1_value = tl.reshape(
e2m1_value, [BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2]
)
evens, odds = tl.split(e2m1_value)
x_fp4 = evens | (odds << 4)
x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
import aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 as _kernel_module
_kernel_module._mxfp4_quant_op = _mxfp4_quant_op_asm_exact
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
_bf16 = dtypes.bf16
_fp4x2 = dtypes.fp4x2
_fp8_e8m0 = dtypes.fp8_e8m0
_kernel_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_ASM_SPLITK = {
(64, 7168, 2048): 2,
(256, 3072, 1536): 1,
}
_FUSED_CONFIGS = {
# M=4,N=2880: BLOCK_SIZE_N=64 for exact N fit (45 tiles), num_stages=2
(4, 2880, 512): {
"BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
"waves_per_eu": 2, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": 1,
},
# M=16: BLOCK_SIZE_N=64 (exact fit for N=2112), num_stages=2, KSPLIT=7
(16, 2112, 7168): {
"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
"waves_per_eu": 1, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": 7,
},
}
_QUANT_BLOCK = 32
_QUANT_TILE = 128
_bufs = {}
def _get_asm_bufs(m, k, n, device):
x_fp4 = torch.empty((m, k >> 1), dtype=torch.uint8, device=device)
sN = (k + _QUANT_BLOCK - 1) // _QUANT_BLOCK
sN_pad = ((sN + 7) >> 3) << 3
sM_pad = ((m + 255) >> 8) << 8
scale = torch.empty((sM_pad, sN_pad), dtype=torch.uint8, device=device)
padded_m = ((m + 31) >> 5) << 5
out = torch.empty((padded_m, n), dtype=_bf16, device=device)
return x_fp4, scale, sN, sN_pad, sM_pad, out, padded_m
@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
n = B.shape[0]
key = (m, n, k)
if key in _ASM_SPLITK:
if key not in _bufs:
_bufs[key] = ('asm', _get_asm_bufs(m, k, n, A.device))
_, (x_fp4, scale, sN, sN_pad, sM_pad, out, padded_m) = _bufs[key]
grid = ((m + _QUANT_TILE - 1) // _QUANT_TILE, sN_pad)
_dynamic_mxfp4_quant_kernel_asm_layout[grid](
A, x_fp4, scale,
A.stride(0), A.stride(1),
x_fp4.stride(0), x_fp4.stride(1),
scale.stride(0), scale.stride(1),
M=m, N=k, scaleN=sN,
scaleM_pad=sM_pad, scaleN_pad=sN_pad,
BLOCK_SIZE=_QUANT_TILE,
MXFP4_QUANT_BLOCK_SIZE=_QUANT_BLOCK,
SCALING_MODE=0, SHUFFLE=True,
)
splitK = _ASM_SPLITK[key]
gemm_a4w4_asm(
x_fp4.view(_fp4x2), B_shuffle, scale.view(_fp8_e8m0), B_scale_sh,
out, _kernel_32x128,
bpreshuffle=True, log2_k_split=splitK,
)
return out[:m]
else:
if key not in _bufs:
_bufs[key] = ('fused', torch.empty((m, n), dtype=torch.bfloat16, device=A.device))
_, out = _bufs[key]
w = B_shuffle.view(torch.uint8).reshape(n // 16, k // 2 * 16)
sm, sn = B_scale_sh.shape
w_scales = B_scale_sh.view(torch.uint8).reshape(sm // 32, sn * 32)
config = _FUSED_CONFIGS.get(key)
return gemm_a16wfp4_preshuffle(A, w, w_scales, prequant=True, y=out, config=config)
scrolls · 153 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 531035.
⋯ 1 unchanged lines#!POPCORN gpu MI355X"""- Version 119: v103 base + ONLY remap_xcd optimization.- Isolate the effect of XCD-aware scheduling without .wt store or acc pattern.+ Version 103: Best combined - v99 + v101's M=4 BLOCK_SIZE_N=64 improvement.+ - M=4: BLOCK_SIZE_N=64 (exact fit for N=2880), num_stages=2 → 6.64µs (was 7.24)+ - M=16: BLOCK_SIZE_N=64, num_stages=2, KSPLIT=7 → 14.2µs+ - M=32: Default fused (no explicit config) → 9.5µs+ - M=64: ASM splitK=2 → 13.7µs+ - M=256: ASM splitK=1 → 12.5µs+ Expected geomean: ~10.5µs"""import torchimport triton⋯ 2 unchanged linesfrom aiter import dtypesfrom aiter.ops.gemm_op_a4w4 import gemm_a4w4_asmfrom aiter.utility.fp4_utils import _dynamic_mxfp4_quant_kernel_asm_layout- from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid, remap_xcdfrom task import input_t, output_t⋯ 35 unchanged linesreturn x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)- _mxfp4_quant_op = _mxfp4_quant_op_asm_exact-import aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 as _kernel_module_kernel_module._mxfp4_quant_op = _mxfp4_quant_op_asm_exact-- # Kernel with ONLY remap_xcd (no .wt store, no acc pattern change)- @triton.heuristics(- {- "EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)- and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)- and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),- "GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])- * triton.cdiv(args["N"], args["BLOCK_SIZE_N"]),- }- )- @triton.jit- def _gemm_a16wfp4_preshuffle_kernel_xcd(- a_ptr, b_ptr, c_ptr, b_scales_ptr,- M, N, K,- stride_am, stride_ak,- stride_bn, stride_bk,- stride_ck, stride_cm, stride_cn,- stride_bsn, stride_bsk,- BLOCK_SIZE_M: tl.constexpr,- BLOCK_SIZE_N: tl.constexpr,- BLOCK_SIZE_K: tl.constexpr,- GROUP_SIZE_M: tl.constexpr,- NUM_KSPLIT: tl.constexpr,- SPLITK_BLOCK_SIZE: tl.constexpr,- EVEN_K: tl.constexpr,- num_warps: tl.constexpr,- num_stages: tl.constexpr,- waves_per_eu: tl.constexpr,- matrix_instr_nonkdim: tl.constexpr,- GRID_MN: tl.constexpr,- PREQUANT: tl.constexpr,- cache_modifier: tl.constexpr,- ):- tl.assume(stride_am > 0)- tl.assume(stride_ak > 0)- tl.assume(stride_bk > 0)- tl.assume(stride_bn > 0)- tl.assume(stride_cm > 0)- tl.assume(stride_cn > 0)- tl.assume(stride_bsk > 0)- tl.assume(stride_bsn > 0)-- pid_unified = tl.program_id(axis=0)- # ONLY change: XCD-aware remapping- pid_unified = remap_xcd(pid_unified, GRID_MN * NUM_KSPLIT, NUM_XCDS=8)-- pid_k = pid_unified % NUM_KSPLIT- pid = pid_unified // NUM_KSPLIT- num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)- num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)-- if NUM_KSPLIT == 1:- pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)- else:- pid_m = pid // num_pid_n- pid_n = pid % num_pid_n-- tl.assume(pid_m >= 0)- tl.assume(pid_n >= 0)- tl.assume(pid_k >= 0)-- SCALE_GROUP_SIZE: tl.constexpr = 32-- if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:- num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)-- offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)- offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16- offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M- a_ptrs = a_ptr + (- offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak- )-- offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)- offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr- offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N- b_ptrs = b_ptr + (- offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk- )-- offs_bsn = (- pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, (BLOCK_SIZE_N // 32))- ) % N- offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(- 0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32- )- b_scale_ptrs = (- b_scales_ptr- + offs_bsn[:, None] * stride_bsn- + offs_ks[None, :] * stride_bsk- )-- accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)-- for k in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):- b_scales = (- tl.load(b_scale_ptrs, cache_modifier=cache_modifier)- .reshape(- BLOCK_SIZE_N // 32,- BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8,- 4, 16, 2, 2, 1,- )- .permute(0, 5, 3, 1, 4, 2, 6)- .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)- )-- if EVEN_K:- a_bf16 = tl.load(a_ptrs)- b = tl.load(b_ptrs, cache_modifier=cache_modifier)-- b = (- b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)- .permute(0, 1, 4, 2, 3, 5)- .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)- .trans(1, 0)- )-- if PREQUANT:- a, a_scales = _mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)-- # Keep original += pattern- accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")-- a_ptrs += BLOCK_SIZE_K * stride_ak- b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk- b_scale_ptrs += BLOCK_SIZE_K * stride_bsk-- c = accumulator.to(c_ptr.type.element_ty)-- offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)- offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)- c_ptrs = (- c_ptr- + stride_cm * offs_cm[:, None]- + stride_cn * offs_cn[None, :]- + pid_k * stride_ck- )- c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)- # Keep original store (no .wt)- tl.store(c_ptrs, c, mask=c_mask)--- import aiter.ops.triton.gemm.basic.gemm_a16wfp4 as _wrapper_module- _wrapper_module._gemm_a16wfp4_preshuffle_kernel = _gemm_a16wfp4_preshuffle_kernel_xcd-from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle_bf16 = dtypes.bf16⋯ 8 unchanged lines}_FUSED_CONFIGS = {+ # M=4,N=2880: BLOCK_SIZE_N=64 for exact N fit (45 tiles), num_stages=2(4, 2880, 512): {"BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 512,"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,"waves_per_eu": 2, "matrix_instr_nonkdim": 16,"cache_modifier": ".cg", "NUM_KSPLIT": 1,},+ # M=16: BLOCK_SIZE_N=64 (exact fit for N=2112), num_stages=2, KSPLIT=7(16, 2112, 7168): {"BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 512,"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
scrolls · 197 diff lines total
Best evidence level for this revision: reported
JSON