submission 755042
musicofhel · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 109 lines, June 9 Researcher Reciprocity License v1.0.
gemm_v4.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-755042?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:2b622658d7d671e365faca9a1620bf45df357efdc175e0bf21efed25298db9bd
license declaredunknown
license concludedunknown
authorsmusicofhel
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
split-k
GEMM v4: Shape-specific ASM kernel tile selection + splitK.Kernel source
gemm_v4.py109 lines
"""
GEMM v4: Shape-specific ASM kernel tile selection + splitK.
Forces optimal tile size AND splitK per benchmark shape.
Available .co tiles: 32x{128-1024}, 64x{128-1024}, 96x128,
128x{128-512}, 160x{128-384}, 192x{128-256}, 224-256x{128-256}
Strategy:
- tile_M >= M (avoid wasted rows from padding to tile boundary)
- tile_N chosen to create ~256/splitK total tiles for good CU util
- splitK fills remaining idle CUs along K dimension
"""
import torch
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
try:
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
except ImportError:
try:
from aiter import gemm_a4w4_asm
except ImportError:
gemm_a4w4_asm = None
_dyn_quant = dynamic_mxfp4_quant
_shuffle = e8m0_shuffle
_asm_gemm = gemm_a4w4_asm
_fallback_gemm = aiter.gemm_a4w4
_fp4x2 = dtypes.fp4x2
_e8m0 = dtypes.fp8_e8m0
def _mangle(tile_m, tile_n):
"""Construct C++ mangled kernel name for given tile size."""
name = f"f4gemm_bf16_per1x32Fp4_BpreShuffle_{tile_m}x{tile_n}"
return f"_ZN5aiter{len(name)}{name}E"
# Pre-computed optimal configs per (M, N, K)
# Format: (kernelName, log2_k_split)
# tile_M: 32 for M<=32, 64 for M<=64, 128/192/256 for larger M
# splitK: fill idle CUs
_SHAPE_CONFIGS = {
# splitK=0 for all (splitK causes accuracy failures with FP4!)
# M=4: 32x128 is the smallest M-tile available
(4, 2880, 512): (_mangle(32, 128), 0),
# M=16: 32x128 (exact tile match not available)
(16, 2112, 7168): (_mangle(32, 128), 0),
# M=32: 32x128 (exact M match)
(32, 4096, 512): (_mangle(32, 128), 0),
(32, 2880, 512): (_mangle(32, 128), 0),
# M=64: force 64x128 (exact M match, less waste than 192x128)
(64, 7168, 2048): (_mangle(64, 128), 0),
# M=256: force 256x128 (exact M match, less waste than 192x128)
(256, 3072, 1536): (_mangle(256, 128), 0),
# Test shapes:
(8, 2112, 7168): (_mangle(32, 128), 0),
(16, 3072, 1536): (_mangle(32, 128), 0),
(64, 3072, 1536): (_mangle(64, 128), 0),
(256, 2880, 512): (_mangle(256, 128), 0),
}
# Default fallback config
_DEFAULT_SPLITK = 0
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]
# Quantize A to MXFP4
A_q, A_scale = _dyn_quant(A.contiguous())
A_scale = _shuffle(A_scale)
A_q_fp4 = A_q.view(_fp4x2)
A_scale_e8m0 = A_scale.view(_e8m0)
if _asm_gemm is None:
return _fallback_gemm(A_q_fp4, B_shuffle, A_scale_e8m0, B_scale_sh,
dtype=torch.bfloat16, bpreshuffle=True)
# Get shape-specific config
config = _SHAPE_CONFIGS.get((m, n, k))
if config is not None:
kernel_name, splitK = config
else:
kernel_name, splitK = "", _DEFAULT_SPLITK
padded_m = (m + 31) // 32 * 32
out = torch.empty((padded_m, n), dtype=torch.bfloat16, device='cuda')
_asm_gemm(
A_q_fp4,
B_shuffle,
A_scale_e8m0,
B_scale_sh,
out,
kernel_name,
None, # bias
1.0, # alpha
0.0, # beta
True, # bpreshuffle
splitK,
)
return out[:m]
scrolls · 109 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