submission 733679
Behzod12312121 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 231 lines, June 9 Researcher Reciprocity License v1.0.
submission_2880_32x128_s3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-733679?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:5329eedc53dc5d967d993b0f2f4dc32678314184c515e415672ac183970349e9
license declaredunknown
license concludedunknown
authorsBehzod12312121
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
split-k
kernel as the working 2112 champion, with splitK=3 baked in.tile-m = 1
BLOCK_M: tl.constexpr = 1,Kernel source
submission_2880_32x128_s3.py231 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
GEMM exact-shape route override for the 2880 family using the same 32x128 ASM
kernel as the working 2112 champion, with splitK=3 baked in.
Shapes overridden:
- (16, 2112, 7168) -> keep the known-good splitK=21 route
- (4, 2880, 512) -> test 32x128 kernel with splitK from env
- (32, 2880, 512) -> test 32x128 kernel with splitK from env
"""
import json
import os
import sys
import torch
import triton
import triton.language as tl
from task import input_t, output_t
_BENCHMARK_SHAPES = {
(4, 2880, 512),
(16, 2112, 7168),
(32, 4096, 512),
(32, 2880, 512),
(64, 7168, 2048),
(256, 3072, 1536),
}
_TARGET_2112 = (16, 2112, 7168)
_TARGET_2880_SHAPES = {
(4, 2880, 512),
(32, 2880, 512),
}
_KERNEL_32X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_SPLIT_K_2880 = 3
_TRACE_SEEN: set[tuple[int, int, int]] = set()
_PATCHED = False
def _emit_route(shape: tuple[int, int, int], config: dict | None, source: str) -> None:
route = "asm"
kernel_name = ""
split_k = None
if config is not None:
kernel_name = str(config.get("kernelName", ""))
split_k = config.get("splitK")
if "_ZN" not in kernel_name:
route = "blockscale"
payload = {
"shape": list(shape),
"ck_config_found": config is not None,
"kernelName": kernel_name,
"splitK": split_k,
"route": route,
"source": source,
}
print(f"GEMM_TRACE {json.dumps(payload, sort_keys=True)}", file=sys.stderr)
def _install_route_patch() -> None:
global _PATCHED
if _PATCHED:
return
import aiter.ops.gemm_op_a4w4 as gemm_mod
orig_get_config = gemm_mod.get_GEMM_config
def wrapped_get_config(m: int, n: int, k: int):
shape = (m, n, k)
if shape == _TARGET_2112:
config = {
"kernelName": _KERNEL_32X128,
"splitK": 21,
}
if shape not in _TRACE_SEEN:
_emit_route(shape, config, source="manual_override_2112")
_TRACE_SEEN.add(shape)
return config
if shape in _TARGET_2880_SHAPES:
config = {
"kernelName": _KERNEL_32X128,
"splitK": _SPLIT_K_2880,
}
if shape not in _TRACE_SEEN:
_emit_route(shape, config, source=f"manual_override_2880_32x128_s{_SPLIT_K_2880}")
_TRACE_SEEN.add(shape)
return config
config = orig_get_config(m, n, k)
if shape in _BENCHMARK_SHAPES and shape not in _TRACE_SEEN:
_emit_route(shape, config, source="baseline")
_TRACE_SEEN.add(shape)
return config
gemm_mod.get_GEMM_config = wrapped_get_config
_PATCHED = True
@triton.jit
def _mxfp4_quant_kernel(
x_ptr,
out_ptr,
scale_ptr,
M,
K,
stride_xm,
GROUP_SIZE: tl.constexpr = 32,
BLOCK_M: tl.constexpr = 1,
):
row = tl.program_id(0)
group_id = tl.program_id(1)
k_start = group_id * GROUP_SIZE
half_offs = tl.arange(0, GROUP_SIZE // 2)
k_even = k_start + half_offs * 2
k_odd = k_start + half_offs * 2 + 1
mask_e = k_even < K
mask_o = k_odd < K
x_even = tl.load(x_ptr + row * stride_xm + k_even, mask=mask_e, other=0.0).to(
tl.float32
)
x_odd = tl.load(x_ptr + row * stride_xm + k_odd, mask=mask_o, other=0.0).to(
tl.float32
)
abs_max = tl.maximum(
tl.max(tl.abs(x_even), axis=0), tl.max(tl.abs(x_odd), axis=0)
)
abs_max = tl.maximum(abs_max, 1e-38).to(tl.float32)
abs_max_int = abs_max.to(tl.int32, bitcast=True)
abs_max_rounded = (
(abs_max_int + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
).to(tl.float32, bitcast=True)
scale_e8m0_unbiased = tl.floor(tl.math.log2(abs_max_rounded)).to(tl.int32) - 2
scale_e8m0_unbiased = tl.minimum(
tl.maximum(scale_e8m0_unbiased, -127), 127
)
e8m0_exp = (scale_e8m0_unbiased + 127).to(tl.uint8)
quant_scale = tl.math.exp2(-scale_e8m0_unbiased.to(tl.float32))
tl.store(scale_ptr + row * (K // GROUP_SIZE) + group_id, e8m0_exp)
xs_e = x_even * quant_scale
xs_e_uint = xs_e.to(tl.int32, bitcast=True).to(tl.uint32)
s_e = xs_e_uint & 0x80000000
xs_e_pos_uint = xs_e_uint ^ s_e
xs_e_pos = xs_e_pos_uint.to(tl.float32, bitcast=True)
sat_e = xs_e_pos >= 6.0
den_e = xs_e_pos < 1.0
mant_odd_e = (xs_e_pos_uint >> 22) & 1
norm_e = (
(xs_e_pos_uint.to(tl.int32) + (-1054867457)) + mant_odd_e.to(tl.int32)
) >> 22
norm_e = norm_e.to(tl.uint8)
den_val_e = (xs_e_pos + 4194304.0).to(tl.int32, bitcast=True) - 0x4A800000
den_val_e = den_val_e.to(tl.uint8)
q_e = tl.full(xs_e.shape, 7, dtype=tl.uint8)
q_e = tl.where(~sat_e, norm_e, q_e)
q_e = tl.where(den_e, den_val_e, q_e)
sign_e_lp = (s_e >> 28).to(tl.uint8)
lo = (q_e | sign_e_lp) & 0xF
xs_o = x_odd * quant_scale
xs_o_uint = xs_o.to(tl.int32, bitcast=True).to(tl.uint32)
s_o = xs_o_uint & 0x80000000
xs_o_pos_uint = xs_o_uint ^ s_o
xs_o_pos = xs_o_pos_uint.to(tl.float32, bitcast=True)
sat_o = xs_o_pos >= 6.0
den_o = xs_o_pos < 1.0
mant_odd_o = (xs_o_pos_uint >> 22) & 1
norm_o = (
(xs_o_pos_uint.to(tl.int32) + (-1054867457)) + mant_odd_o.to(tl.int32)
) >> 22
norm_o = norm_o.to(tl.uint8)
den_val_o = (xs_o_pos + 4194304.0).to(tl.int32, bitcast=True) - 0x4A800000
den_val_o = den_val_o.to(tl.uint8)
q_o = tl.full(xs_o.shape, 7, dtype=tl.uint8)
q_o = tl.where(~sat_o, norm_o, q_o)
q_o = tl.where(den_o, den_val_o, q_o)
sign_o_lp = (s_o >> 28).to(tl.uint8)
hi = ((q_o | sign_o_lp) & 0xF) << 4
packed = lo | hi
out_base = row * (K // 2) + k_start // 2
tl.store(out_ptr + out_base + half_offs, packed.to(tl.uint8))
def _triton_mxfp4_quant(x: torch.Tensor):
m, k = x.shape
assert k % 32 == 0
x = x.contiguous()
out = torch.empty(m, k // 2, dtype=torch.uint8, device=x.device)
scale = torch.empty(m, k // 32, dtype=torch.uint8, device=x.device)
grid = (m, k // 32)
_mxfp4_quant_kernel[grid](x, out, scale, m, k, x.stride(0))
return out, scale
def custom_kernel(data: input_t) -> output_t:
import aiter
from aiter import dtypes
from aiter.utility.fp4_utils import e8m0_shuffle
_install_route_patch()
a, _, _, b_shuffle, b_scale_sh = data
a = a.contiguous()
a_fp4, a_scale = _triton_mxfp4_quant(a)
a_scale_sh = e8m0_shuffle(a_scale.view(torch.uint8))
return aiter.gemm_a4w4(
a_fp4.view(dtypes.fp4x2),
b_shuffle,
a_scale_sh.view(dtypes.fp8_e8m0),
b_scale_sh,
dtype=dtypes.bf16,
bpreshuffle=True,
)
scrolls · 231 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