submission 671354
sizezheng_94252 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 127 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-671354?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:9488cdf75e3b97bca304536bb7e0ec9d084798d9d6d1d321b8953a6ce327d9b8
license declaredunknown
license concludedunknown
authorssizezheng_94252
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
Optimized MXFP4 GEMM: gemm_a16wfp4 with fused quant + fast B scale computation.tile-n = 64
BLOCK_N = 64Kernel source
submission.py127 lines
"""
Optimized MXFP4 GEMM: gemm_a16wfp4 with fused quant + fast B scale computation.
Key: eliminates full B re-quantization. Only computes B scales (tiny kernel).
The Triton GEMM kernel handles A quant internally (fused).
"""
from task import input_t, output_t
import torch
import triton
import triton.language as tl
from aiter import dtypes
from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4
from aiter.ops.triton.quant import dynamic_mxfp4_quant
@triton.jit
def _compute_b_scale_kernel(
B_ptr, Scale_ptr,
N, K,
stride_bn, stride_bk,
BLOCK_N: tl.constexpr,
GROUP_SIZE: tl.constexpr,
):
"""Compute e8m0 per-group-of-32 scale from bf16 B. No quantization of data."""
pid = tl.program_id(0)
n_start = pid * BLOCK_N
n_offsets = n_start + tl.arange(0, BLOCK_N)
n_mask = n_offsets < N
num_groups = K // GROUP_SIZE
for g in range(num_groups):
k_start = g * GROUP_SIZE
# Load bf16 group and find abs max
amax = tl.zeros([BLOCK_N], dtype=tl.float32)
for k_off in range(0, GROUP_SIZE, 16):
k_offsets = k_start + k_off + tl.arange(0, 16)
ptrs = B_ptr + n_offsets[:, None] * stride_bn + k_offsets[None, :]
mask = n_mask[:, None] & (k_offsets[None, :] < K)
vals = tl.load(ptrs, mask=mask, other=0.0).to(tl.float32)
amax = tl.maximum(amax, tl.max(tl.abs(vals), axis=1))
# Match aiter's exact formula:
# 1. Round amax UP to nearest power of 2 (via float32 bit manipulation)
# 2. scale_unbiased = floor(log2(amax_rounded)) - 2
# 3. e8m0 = scale_unbiased + 127
amax = tl.maximum(amax, 1e-12)
# Round up to power of 2: add 0x200000 (half ULP of mantissa) then mask
amax_i32 = amax.to(tl.int32, bitcast=True)
amax_i32 = (amax_i32 + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax_rounded = amax_i32.to(tl.float32, bitcast=True)
scale_unbiased = tl.math.floor(tl.math.log2(amax_rounded)) - 2
scale_unbiased = tl.maximum(tl.minimum(scale_unbiased, 127), -127)
e8m0 = (scale_unbiased.to(tl.int32) + 127).to(tl.uint8)
# Store: row-major (N, K//32)
s_ptrs = Scale_ptr + n_offsets * num_groups + g
tl.store(s_ptrs, e8m0, mask=n_mask)
def compute_b_scale(B, K):
"""Compute e8m0 scales from bf16 B. Much faster than full quant."""
N = B.shape[0]
num_groups = K // 32
scale = torch.empty((N, num_groups), dtype=torch.uint8, device=B.device)
BLOCK_N = 64
grid = ((N + BLOCK_N - 1) // BLOCK_N,)
_compute_b_scale_kernel[grid](
B, scale, N, K,
B.stride(0), B.stride(1),
BLOCK_N=BLOCK_N, GROUP_SIZE=32,
)
return scale
# Per-shape optimal configs
_SHAPE_CONFIGS = {
(4, 2880, 512): {'BLOCK_SIZE_M': 4, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 256,
'GROUP_SIZE_M': 1, 'num_warps': 4, 'num_stages': 2,
'waves_per_eu': 3, 'matrix_instr_nonkdim': 16,
'cache_modifier': None, 'NUM_KSPLIT': 1},
(16, 2112, 7168): {'BLOCK_SIZE_M': 4, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 256,
'GROUP_SIZE_M': 1, 'num_warps': 4, 'num_stages': 2,
'waves_per_eu': 3, 'matrix_instr_nonkdim': 16,
'cache_modifier': None, 'NUM_KSPLIT': 16},
(32, 4096, 512): {'BLOCK_SIZE_M': 4, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 512,
'GROUP_SIZE_M': 1, 'num_warps': 4, 'num_stages': 2,
'waves_per_eu': 2, 'matrix_instr_nonkdim': 16,
'cache_modifier': None, 'NUM_KSPLIT': 1},
(32, 2880, 512): {'BLOCK_SIZE_M': 4, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 512,
'GROUP_SIZE_M': 1, 'num_warps': 4, 'num_stages': 2,
'waves_per_eu': 2, 'matrix_instr_nonkdim': 16,
'cache_modifier': None, 'NUM_KSPLIT': 1},
(64, 7168, 2048): {'BLOCK_SIZE_M': 16, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 256,
'GROUP_SIZE_M': 1, 'num_warps': 4, 'num_stages': 2,
'waves_per_eu': 2, 'matrix_instr_nonkdim': 16,
'cache_modifier': None, 'NUM_KSPLIT': 1},
(256, 3072, 1536): {'BLOCK_SIZE_M': 16, 'BLOCK_SIZE_N': 256, 'BLOCK_SIZE_K': 256,
'GROUP_SIZE_M': 1, 'num_warps': 4, 'num_stages': 2,
'waves_per_eu': 2, 'matrix_instr_nonkdim': 16,
'cache_modifier': None, 'NUM_KSPLIT': 1},
}
_DEFAULT_CFG = {'BLOCK_SIZE_M': 4, 'BLOCK_SIZE_N': 128, 'BLOCK_SIZE_K': 256,
'GROUP_SIZE_M': 1, 'num_warps': 4, 'num_stages': 2,
'waves_per_eu': 3, 'matrix_instr_nonkdim': 16,
'cache_modifier': None, 'NUM_KSPLIT': 1}
_b_cache = {}
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
m, k = A.shape
n = B.shape[0]
# Cache B_scale
b_ptr = B.data_ptr()
if b_ptr not in _b_cache:
_b_cache.clear()
_, B_scale = dynamic_mxfp4_quant(B.contiguous())
_b_cache[b_ptr] = (B_scale, B)
B_scale, _ = _b_cache[b_ptr]
B_q_u8 = B_q.view(torch.uint8)
cfg = _SHAPE_CONFIGS.get((m, n, k), _DEFAULT_CFG)
return gemm_a16wfp4(A, B_q_u8, B_scale, False, dtype=dtypes.bf16, config=cfg)
scrolls · 127 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