submission 715401
Hamza · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 563 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-715401?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:17f5a15763d609cd75ffbdbcce2428889d18314775025825fc909efa26751219
license declaredunknown
license concludedunknown
authorsHamza
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 4
num_warps=4, num_stages=2, waves_per_eu=WPE,split-k
_lines = ["cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio"]stages = 2
num_warps=4, num_stages=2, waves_per_eu=WPE,tile-k = 256
BLOCK_K = 256 if K_real <= KSPLIT * 512 or (KSPLIT == 2 and K_real <= KSPLIT * 1024) else 512tile-m = 8
BLOCK_M = 8tile-n = 128
BLOCK_N = 128vector-width = float4
float4 s = *reinterpret_cast<const float4*>(pp + idx4);Kernel source
submission.py563 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
# submission_v19_fastquant.py — Optimized E8M0 scale computation (integer bit ops)
# Replace tl.log2().floor() + tl.exp2() with integer shift/mask in _mxfp4_quant_op
# --- Config injection (prevents extra module_gemm_common build ~20s) ---
import os as _os
# Must be set BEFORE torch import for load_inline HIP compilation
_os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
_os.environ.setdefault("CXX", "clang++")
# Fresh Triton cache to force recompilation with modified quant source
import uuid as _uuid
_os.environ["TRITON_CACHE_DIR"] = f"/tmp/_triton_fq_{_uuid.uuid4().hex[:8]}"
_KERNEL_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_CSV_PATH = "/tmp/_mxfp4_mm_config.csv"
_CU = 256
_NK_FAMILIES = [
(2880, 512), (2112, 7168), (4096, 512), (7168, 2048), (3072, 1536),
(2880, 1536), (4096, 1536), (2112, 512), (2112, 2048),
(7168, 512), (7168, 1536), (7168, 7168), (3072, 512),
(3072, 7168), (3072, 2048), (4096, 2048), (4096, 7168),
(2880, 2048), (2880, 7168),
]
_M_VALUES = [1, 2, 4, 8, 16, 32, 64, 128, 256]
_lines = ["cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio"]
for _n, _k in _NK_FAMILIES:
for _m in _M_VALUES:
_tile_num = ((_m + 31) // 32) * ((_n + 127) // 128)
_cus_per_tile = _CU / max(_tile_num, 1)
_split = 0
while _cus_per_tile >= pow(2, _split + 1) and (pow(2, _split + 1) * 128) < 2 * _k:
_split += 1
_split = min(_split, 3)
_lines.append(f"{_CU},{_m},{_n},{_k},21,{_split},1.0,{_KERNEL_32x128},0,0,0.0")
with open(_CSV_PATH, "w") as _f:
_f.write("\n".join(_lines))
_os.environ["AITER_CONFIG_GEMM_A4W4"] = _CSV_PATH + ":/home/runner/aiter/aiter/configs/a4w4_blockscale_tuned_gemm.csv"
# --- End config injection ---
import torch
torch.set_grad_enabled(False)
import triton
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (
_gemm_a16wfp4_preshuffle_kernel,
)
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
_gemm_afp4wfp4_reduce_kernel,
)
from task import input_t, output_t
import sys as _sys
import time as _time
import gc as _gc
_sys.setswitchinterval(1.0)
# --- Monkey-patch heuristics to constants ---
try:
_gemm_a16wfp4_preshuffle_kernel.values['GRID_MN'] = lambda args: 1
_gemm_a16wfp4_preshuffle_kernel.values['EVEN_K'] = lambda args: True
print("[patch] GRID_MN → 1, EVEN_K → True", file=_sys.stderr, flush=True)
except (AttributeError, KeyError, TypeError) as _e:
print(f"[patch] heuristics failed: {_e}", file=_sys.stderr, flush=True)
_os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
# --- Modify _mxfp4_quant_op: replace log2/floor/exp2 with integer bit ops ---
# The E8M0 scale computation uses tl.log2(amax).floor() - 2, but amax is already
# a power of 2 (mantissa zeroed by & 0xFF800000). So the exponent can be extracted
# with integer bit shifts, eliminating expensive v_log_f32 and v_ldexp_f32 SFU instructions.
print("[fastquant] Modifying _mxfp4_quant_op source...", file=_sys.stderr, flush=True)
try:
_jit_fn = _gemm_a16wfp4_preshuffle_kernel.fn if hasattr(_gemm_a16wfp4_preshuffle_kernel, 'fn') else _gemm_a16wfp4_preshuffle_kernel
_quant_fn = _jit_fn.__globals__['_mxfp4_quant_op']
# _mxfp4_quant_op is a JITFunction — do NOT unwrap via .fn
_old_qsrc = _quant_fn._src
# Replacement 1: log2(amax).floor() → integer bit extraction
_new_qsrc = _old_qsrc.replace(
" amax = amax.to(tl.float32, bitcast=True)\n"
" scale_e8m0_unbiased = tl.log2(amax).floor() - 2",
" amax_exp = (amax >> 23) & 0xFF\n"
" scale_e8m0_unbiased = (amax_exp.to(tl.int32) - 129).to(tl.float32)"
)
# Replacement 2: exp2(-scale) → integer FP32 construction
_new_qsrc = _new_qsrc.replace(
" quant_scale = tl.exp2(-scale_e8m0_unbiased)",
" quant_scale = (((127.0 - scale_e8m0_unbiased).to(tl.int32).to(tl.uint32) << 23)).to(tl.float32, bitcast=True)"
)
if _new_qsrc != _old_qsrc:
if hasattr(_quant_fn, '_unsafe_update_src'):
_quant_fn._unsafe_update_src(_new_qsrc)
else:
_quant_fn._src = _new_qsrc
if hasattr(_quant_fn, 'src'):
_quant_fn.src = _new_qsrc
if hasattr(_quant_fn, 'hash'):
_quant_fn.hash = None
# Also modify the KERNEL source to bust its Triton cache key
# (quant is a dependency but its hash isn't in the kernel's cache key)
_old_ksrc = _jit_fn._src
_new_ksrc = _old_ksrc.replace(
'accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")',
'accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", acc=accumulator)'
)
if _new_ksrc != _old_ksrc:
_jit_fn._unsafe_update_src(_new_ksrc)
print("[fastquant] Applied quant + kernel source modifications", file=_sys.stderr, flush=True)
else:
print("[fastquant] Applied quant mod, kernel mod FAILED", file=_sys.stderr, flush=True)
# Verify
_vq = _quant_fn._src if hasattr(_quant_fn, '_src') else ''
_vk = _jit_fn._src if hasattr(_jit_fn, '_src') else ''
print(f"[fastquant] quant has amax_exp: {'amax_exp' in _vq}, kernel has acc=: {'acc=accumulator' in _vk}",
file=_sys.stderr, flush=True)
else:
print("[fastquant] WARNING: replacement strings not found — source unchanged", file=_sys.stderr, flush=True)
except Exception as _e:
import traceback
print(f"[fastquant] FAILED: {_e}", file=_sys.stderr, flush=True)
traceback.print_exc(file=_sys.stderr)
# --- End quant modification ---
# --- HIP reduce kernel ---
_HIP_REDUCE_SRC = r"""
#include <hip/hip_runtime.h>
__device__ __forceinline__ unsigned short f32_to_bf16(float f) {
unsigned int u;
__builtin_memcpy(&u, &f, sizeof(u));
unsigned int rounding_bias = ((u >> 16) & 1) + 0x7FFFu;
return (unsigned short)((u + rounding_bias) >> 16);
}
template <int KSPLIT>
__global__ void reduce_k_vec4(const float* __restrict__ pp,
unsigned short* __restrict__ out, int MN) {
int idx4 = (blockIdx.x * blockDim.x + threadIdx.x) * 4;
if (idx4 + 3 < MN) {
float4 s = *reinterpret_cast<const float4*>(pp + idx4);
#pragma unroll
for (int k = 1; k < KSPLIT; k++) {
float4 v = *reinterpret_cast<const float4*>(pp + k * MN + idx4);
s.x += v.x; s.y += v.y; s.z += v.z; s.w += v.w;
}
unsigned short r0 = f32_to_bf16(s.x);
unsigned short r1 = f32_to_bf16(s.y);
unsigned short r2 = f32_to_bf16(s.z);
unsigned short r3 = f32_to_bf16(s.w);
*reinterpret_cast<unsigned long long*>(out + idx4) =
(unsigned long long)r0 | ((unsigned long long)r1 << 16) |
((unsigned long long)r2 << 32) | ((unsigned long long)r3 << 48);
} else {
for (int i = idx4; i < MN && i < idx4 + 4; i++) {
float s = pp[i];
#pragma unroll
for (int k = 1; k < KSPLIT; k++) s += pp[k * MN + i];
out[i] = f32_to_bf16(s);
}
}
}
__global__ void reduce_k_gen(const float* __restrict__ pp,
unsigned short* __restrict__ out, int MN, int ksplit) {
int idx = blockIdx.x * blockDim.x + threadIdx.x;
if (idx < MN) {
float s = pp[idx];
for (int k = 1; k < ksplit; k++) s += pp[k * MN + idx];
out[idx] = f32_to_bf16(s);
}
}
void reduce_op(torch::Tensor pp, torch::Tensor out, int M, int N, int ksplit) {
int MN = M * N;
const float* pp_ptr = pp.data_ptr<float>();
unsigned short* out_ptr = reinterpret_cast<unsigned short*>(out.data_ptr());
const int threads_v = 64;
const int elems_per_block = threads_v * 4;
const int blocks_v = (MN + elems_per_block - 1) / elems_per_block;
switch (ksplit) {
case 2: reduce_k_vec4<2><<<blocks_v, threads_v>>>(pp_ptr, out_ptr, MN); break;
case 3: reduce_k_vec4<3><<<blocks_v, threads_v>>>(pp_ptr, out_ptr, MN); break;
case 4: reduce_k_vec4<4><<<blocks_v, threads_v>>>(pp_ptr, out_ptr, MN); break;
case 7: reduce_k_vec4<7><<<blocks_v, threads_v>>>(pp_ptr, out_ptr, MN); break;
case 8: reduce_k_vec4<8><<<blocks_v, threads_v>>>(pp_ptr, out_ptr, MN); break;
default: {
const int threads = 256;
const int blocks = (MN + threads - 1) / threads;
reduce_k_gen<<<blocks, threads>>>(pp_ptr, out_ptr, MN, ksplit);
break;
}
}
}
"""
_HIP_REDUCE_CPP = "void reduce_op(torch::Tensor pp, torch::Tensor out, int M, int N, int ksplit);"
_USE_HIP_REDUCE = False
try:
from torch.utils.cpp_extension import load_inline as _load_inline
_hip_reduce_t0 = _time.time()
_hip_reduce = _load_inline(
name="mxfp4_reduce_hip",
cpp_sources=[_HIP_REDUCE_CPP],
cuda_sources=[_HIP_REDUCE_SRC],
functions=["reduce_op"],
verbose=False,
extra_cuda_cflags=["--offload-arch=gfx950", "-O3"],
)
_USE_HIP_REDUCE = True
print(f"[hip] reduce kernel compiled in {_time.time()-_hip_reduce_t0:.1f}s",
file=_sys.stderr, flush=True)
except Exception as _e:
print(f"[hip] reduce kernel FAILED (using Triton fallback): {_e}",
file=_sys.stderr, flush=True)
# --- End HIP reduce kernel ---
# --- Helper functions ---
def _get_splitk(K: int, BLOCK_SIZE_K: int, NUM_KSPLIT: int):
SPLITK_BLOCK_SIZE = (
triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K
)
while NUM_KSPLIT > 1 and BLOCK_SIZE_K > 16:
if (
K % (SPLITK_BLOCK_SIZE // 2) == 0
and SPLITK_BLOCK_SIZE % BLOCK_SIZE_K == 0
and K % (BLOCK_SIZE_K // 2) == 0
):
break
elif K % (SPLITK_BLOCK_SIZE // 2) != 0 and NUM_KSPLIT > 1:
NUM_KSPLIT = NUM_KSPLIT // 2
elif SPLITK_BLOCK_SIZE % BLOCK_SIZE_K != 0:
if NUM_KSPLIT > 1:
NUM_KSPLIT = NUM_KSPLIT // 2
elif BLOCK_SIZE_K > 16:
BLOCK_SIZE_K = BLOCK_SIZE_K // 2
elif K % (BLOCK_SIZE_K // 2) != 0 and BLOCK_SIZE_K > 16:
BLOCK_SIZE_K = BLOCK_SIZE_K // 2
else:
break
SPLITK_BLOCK_SIZE = (
triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K
)
return SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT
_CFG_CACHE: dict = {}
def _get_cfg(M: int, N: int, K_real: int):
key = (M, N, K_real)
if key in _CFG_CACHE:
return _CFG_CACHE[key]
K = K_real // 2
if M <= 32:
BLOCK_M = 8
BLOCK_N = 128
tiles_128 = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + 127) // 128)
KSPLIT = 1
if K_real >= 4096:
KSPLIT = 7
elif K_real >= 2048:
if tiles_128 * 2 >= (_CU * 3) // 4 and tiles_128 * 2 <= _CU:
KSPLIT = 2
else:
KSPLIT = 4
elif K_real >= 1536:
if tiles_128 * 2 >= (_CU * 3) // 4 and tiles_128 * 2 <= _CU:
KSPLIT = 2
else:
KSPLIT = 3
BLOCK_K = 256 if K_real <= KSPLIT * 512 or (KSPLIT == 2 and K_real <= KSPLIT * 1024) else 512
if tiles_128 * KSPLIT < (_CU * 3) // 4:
BLOCK_N = 64
wgs = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + BLOCK_N - 1) // BLOCK_N) * KSPLIT
cfg = {
"BLOCK_SIZE_M": BLOCK_M, "BLOCK_SIZE_N": BLOCK_N, "BLOCK_SIZE_K": BLOCK_K,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
"waves_per_eu": 2 if wgs > _CU else 1, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": KSPLIT,
}
else:
BLOCK_M = 16
if M <= 128:
tiles_bm16 = ((M + 15) // 16) * ((N + 127) // 128)
if tiles_bm16 < (_CU * 3) // 4:
BLOCK_M = 8
tiles = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + 127) // 128)
BLOCK_N = 128
KSPLIT = 1
if _CU // 2 <= tiles <= _CU and (K_real >= 7168 or (K_real >= 2048 and BLOCK_M == 8)):
KSPLIT = 2
elif tiles < _CU // 2 and K_real > 512:
if K_real >= 4096:
if tiles * 2 >= _CU:
KSPLIT = 2
else:
KSPLIT = 7
elif K_real >= 2048:
KSPLIT = 2
elif K_real >= 1536:
KSPLIT = 3
BLOCK_K = 256 if K_real <= max(KSPLIT * 4096, 2048) else 512
if tiles * KSPLIT < (_CU * 3) // 4:
BLOCK_N = 64
wgs = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + BLOCK_N - 1) // BLOCK_N) * KSPLIT
cfg = {
"BLOCK_SIZE_M": BLOCK_M, "BLOCK_SIZE_N": BLOCK_N, "BLOCK_SIZE_K": BLOCK_K,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
"waves_per_eu": 2 if wgs > _CU else 1, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": KSPLIT,
}
if cfg["NUM_KSPLIT"] > 1:
SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT = _get_splitk(
K, cfg["BLOCK_SIZE_K"], cfg["NUM_KSPLIT"]
)
cfg["SPLITK_BLOCK_SIZE"] = SPLITK_BLOCK_SIZE
cfg["BLOCK_SIZE_K"] = BLOCK_SIZE_K
cfg["NUM_KSPLIT"] = NUM_KSPLIT
if cfg["BLOCK_SIZE_K"] >= 2 * K:
cfg["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * K)
cfg["SPLITK_BLOCK_SIZE"] = 2 * K
cfg["NUM_KSPLIT"] = 1
cfg["BLOCK_SIZE_N"] = max(cfg["BLOCK_SIZE_N"], 32)
if cfg["NUM_KSPLIT"] == 1:
cfg["SPLITK_BLOCK_SIZE"] = 2 * K
actual_ksplit = None
nk_pow2 = None
if cfg["NUM_KSPLIT"] > 1:
actual_ksplit = triton.cdiv(K, cfg["SPLITK_BLOCK_SIZE"] // 2)
nk_pow2 = triton.next_power_of_2(cfg["NUM_KSPLIT"])
num_m_tiles = triton.cdiv(M, cfg["BLOCK_SIZE_M"])
num_n_tiles = triton.cdiv(N, cfg["BLOCK_SIZE_N"])
total_tiles = num_m_tiles * num_n_tiles
grid_main = (cfg["NUM_KSPLIT"] * total_tiles,)
grid_reduce = None
if cfg["NUM_KSPLIT"] > 1:
grid_reduce = (triton.cdiv(M, 16), triton.cdiv(N, 16))
result = (cfg, actual_ksplit, nk_pow2, grid_main, grid_reduce,
K, cfg["BLOCK_SIZE_M"], cfg["BLOCK_SIZE_N"], cfg["BLOCK_SIZE_K"],
cfg["NUM_KSPLIT"], cfg["SPLITK_BLOCK_SIZE"], cfg["waves_per_eu"])
_CFG_CACHE[key] = result
return result
# --- Nuclear pre-warming ---
_WARMUP_T0 = _time.time()
_PREWARMED_CONFIGS = {}
_NO_LSR = {}
_LSR = {}
_REDUCE = set()
for _nw, _kw in _NK_FAMILIES:
for _mw in _M_VALUES:
_cw, _aw, _nkw, _, _, _, _, _, _, _, _, _ = _get_cfg(_mw, _nw, _kw)
_ck = (_cw["BLOCK_SIZE_M"], _cw["BLOCK_SIZE_N"], _cw["BLOCK_SIZE_K"],
_cw["NUM_KSPLIT"], _cw["SPLITK_BLOCK_SIZE"], _cw["waves_per_eu"])
if _mw <= 32 and _kw >= 1536:
_NO_LSR.setdefault(_ck, True)
else:
_LSR.setdefault(_ck, True)
if _aw is not None:
_REDUCE.add((_aw, _nkw))
for _k in _NO_LSR:
_LSR.pop(_k, None)
print(f"[pre-warm] {len(_NO_LSR)} no-lsr + {len(_LSR)} lsr GEMM, {len(_REDUCE)} reduce configs",
file=_sys.stderr, flush=True)
_wA = torch.zeros(32, 8192, dtype=torch.bfloat16, device="cuda")
_wBw = torch.zeros(16, 65536, dtype=torch.uint8, device="cuda")
_wBs = torch.zeros(16, 65536, dtype=torch.uint8, device="cuda")
_wypp = torch.zeros(16, 32, 256, dtype=torch.float32, device="cuda")
_wy = torch.zeros(32, 256, dtype=torch.bfloat16, device="cuda")
def _pw(bm, bn, bk, ks, spk, wpe):
c = {"BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": bn, "BLOCK_SIZE_K": bk,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
"waves_per_eu": wpe, "matrix_instr_nonkdim": 16,
"cache_modifier": ".cg", "NUM_KSPLIT": ks, "SPLITK_BLOCK_SIZE": spk}
o = _wypp if ks > 1 else _wy
_gemm_a16wfp4_preshuffle_kernel[(max(ks, 1),)](
_wA, _wBw, o, _wBs, bm, bn, spk // 2,
_wA.stride(0), _wA.stride(1), _wBw.stride(0), _wBw.stride(1),
0 if ks <= 1 else _wypp.stride(0),
_wy.stride(0) if ks <= 1 else _wypp.stride(1),
_wy.stride(1) if ks <= 1 else _wypp.stride(2),
_wBs.stride(0), _wBs.stride(1), PREQUANT=True, **c)
# Phase 1: M≤32 K>=1536 without disable-lsr
print("[pre-warm] Phase 1: M≤32 K>=1536 (no disable-lsr)...", file=_sys.stderr, flush=True)
for _ck in sorted(_NO_LSR):
try:
_pw(*_ck)
_PREWARMED_CONFIGS[_ck] = "no-lsr"
print(f" BM={_ck[0]} BN={_ck[1]} BK={_ck[2]} KS={_ck[3]} SPK={_ck[4]} wpe={_ck[5]} ({_time.time()-_WARMUP_T0:.0f}s)",
file=_sys.stderr, flush=True)
except Exception as _e:
print(f" {_ck}: FAIL {_e}", file=_sys.stderr, flush=True)
# Phase 2: set disable-lsr
_os.environ["DISABLE_LLVM_OPT"] = "disable-lsr"
print(f"[pre-warm] Phase 2: DISABLE_LLVM_OPT=disable-lsr set ({_time.time()-_WARMUP_T0:.0f}s)",
file=_sys.stderr, flush=True)
# Phase 3: remaining GEMM configs with disable-lsr
_lsr_list = sorted(_LSR)
print(f"[pre-warm] Phase 3: {len(_lsr_list)} remaining GEMM configs (disable-lsr)...",
file=_sys.stderr, flush=True)
for _idx, _ck in enumerate(_lsr_list):
if _time.time() - _WARMUP_T0 > 200:
print(f" timeout safety — {len(_lsr_list) - _idx} configs skipped",
file=_sys.stderr, flush=True)
break
try:
_pw(*_ck)
_PREWARMED_CONFIGS[_ck] = "lsr"
print(f" BM={_ck[0]} BN={_ck[1]} BK={_ck[2]} KS={_ck[3]} SPK={_ck[4]} wpe={_ck[5]} ({_time.time()-_WARMUP_T0:.0f}s)",
file=_sys.stderr, flush=True)
except Exception as _e:
print(f" {_ck}: FAIL {_e}", file=_sys.stderr, flush=True)
# Phase 4: reduce kernel configs
print(f"[pre-warm] Phase 4: {len(_REDUCE)} reduce configs...", file=_sys.stderr, flush=True)
for _ak, _nk in sorted(_REDUCE):
if _time.time() - _WARMUP_T0 > 230:
print(" timeout safety — remaining reduce configs skipped", file=_sys.stderr, flush=True)
break
try:
_gemm_afp4wfp4_reduce_kernel[(1, 1)](
_wypp, _wy, 16, 16,
_wypp.stride(0), _wypp.stride(1), _wypp.stride(2),
_wy.stride(0), _wy.stride(1), 16, 16, _ak, _nk)
print(f" ksplit={_ak} nk_pow2={_nk} ({_time.time()-_WARMUP_T0:.0f}s)",
file=_sys.stderr, flush=True)
except Exception as _e:
print(f" ksplit={_ak} nk={_nk}: FAIL {_e}", file=_sys.stderr, flush=True)
del _wA, _wBw, _wBs, _wypp, _wy, _pw
del _NO_LSR, _LSR, _REDUCE, _lsr_list
torch.cuda.empty_cache()
print(f"[pre-warm] Done: {len(_PREWARMED_CONFIGS)} GEMM configs in {_time.time()-_WARMUP_T0:.0f}s",
file=_sys.stderr, flush=True)
_gc.disable()
# --- End pre-warming ---
_PRESHUFFLE_CACHE: dict = {}
_OUT_BUF: dict = {}
_YPP_BUF: dict = {}
_LOGGED: set = set()
def _get_preshuffle_b(data):
key = data[3].data_ptr()
if key not in _PRESHUFFLE_CACHE:
N = data[3].shape[0]
K_bytes = data[3].shape[1]
sm, sn = data[4].shape
N_groups = N // 32
B_w = data[3].view(torch.uint8).reshape(N // 16, K_bytes * 16)
B_s = data[4].view(torch.uint8).reshape(sm // 32, sn * 32)[:N_groups].contiguous()
_PRESHUFFLE_CACHE[key] = (B_w, B_s, B_w.stride(0), B_s.stride(0))
return _PRESHUFFLE_CACHE[key]
def custom_kernel(data: input_t) -> output_t:
A = data[0]
if not A.is_contiguous():
A = A.contiguous()
_ndim = A.ndim
if _ndim == 2:
A_2d = A
M = A.shape[0]
else:
A_2d = A.view(-1, A.shape[-1])
M = A_2d.shape[0]
N = data[3].shape[0]
K_bytes = data[3].shape[1]
K_real = K_bytes * 2
cfg, actual_ksplit, nk_pow2, grid_main, grid_reduce, K, BM, BN, BK, KS, SPK, WPE = _get_cfg(M, N, K_real)
_sk = (M, N, K_real)
if _sk not in _LOGGED:
_LOGGED.add(_sk)
print(f"[kernel] M={M} N={N} K={K_real} BM={BM} BN={BN} "
f"BK={BK} KS={KS} wpe={WPE} "
f"grid={grid_main[0]}", file=_sys.stderr, flush=True)
okey = (M, N)
if okey not in _OUT_BUF:
_OUT_BUF[okey] = torch.empty((M, N), dtype=torch.bfloat16, device="cuda")
y = _OUT_BUF[okey]
B_w, B_s, stride_bw0, stride_bs0 = _get_preshuffle_b(data)
if KS > 1:
ppkey = (nk_pow2, M, N)
if ppkey not in _YPP_BUF:
_YPP_BUF[ppkey] = torch.empty(
(nk_pow2, M, N), dtype=torch.float32, device="cuda"
)
y_pp = _YPP_BUF[ppkey]
stride_ck = M * N
stride_cm = N
else:
y_pp = None
stride_ck = 0
stride_cm = N
_gemm_a16wfp4_preshuffle_kernel[grid_main](
A_2d, B_w,
y if y_pp is None else y_pp,
B_s,
M, N, K,
K_real, 1,
stride_bw0, 1,
stride_ck, stride_cm, 1,
stride_bs0, 1,
BLOCK_SIZE_M=BM, BLOCK_SIZE_N=BN, BLOCK_SIZE_K=BK,
GROUP_SIZE_M=1, NUM_KSPLIT=KS, SPLITK_BLOCK_SIZE=SPK,
num_warps=4, num_stages=2, waves_per_eu=WPE,
matrix_instr_nonkdim=16, cache_modifier=".cg",
PREQUANT=True,
)
if y_pp is not None:
if _USE_HIP_REDUCE:
_hip_reduce.reduce_op(y_pp, y, M, N, actual_ksplit)
else:
_gemm_afp4wfp4_reduce_kernel[grid_reduce](
y_pp, y, M, N,
M * N, N, 1,
N, 1,
16, 16,
actual_ksplit, nk_pow2,
)
if _ndim == 2:
return y
return y.view(*A.shape[:-1], N)
scrolls · 563 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 698755.
#!POPCORN leaderboard amd-mxfp4-mm#!POPCORN gpu MI355X- # submission_direct.py v7 — Nuclear pre-warming + selective disable-lsr + HIP_FORCE_DEV_KERNARG+ # submission_v19_fastquant.py — Optimized E8M0 scale computation (integer bit ops)+ # Replace tl.log2().floor() + tl.exp2() with integer shift/mask in _mxfp4_quant_op# --- Config injection (prevents extra module_gemm_common build ~20s) ---import os as _os⋯ 1 unchanged lines# Must be set BEFORE torch import for load_inline HIP compilation_os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")_os.environ.setdefault("CXX", "clang++")+ # Fresh Triton cache to force recompilation with modified quant source+ import uuid as _uuid+ _os.environ["TRITON_CACHE_DIR"] = f"/tmp/_triton_fq_{_uuid.uuid4().hex[:8]}"_KERNEL_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"_CSV_PATH = "/tmp/_mxfp4_mm_config.csv"⋯ 22 unchanged lines# --- End config injection ---import torch+ torch.set_grad_enabled(False)import tritonfrom aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (_gemm_a16wfp4_preshuffle_kernel,⋯ 5 unchanged linesimport sys as _sysimport time as _timeimport gc as _gc+ _sys.setswitchinterval(1.0)# --- Monkey-patch heuristics to constants ---- # GRID_MN: dead tl.constexpr creating separate cache entries per (M,N,BM,BN).- # EVEN_K: always True due to _get_splitk alignment logic. Skip the modulo checks.- # Both patches reduce per-call Python overhead (lambda evaluation) by ~1µs.try:_gemm_a16wfp4_preshuffle_kernel.values['GRID_MN'] = lambda args: 1_gemm_a16wfp4_preshuffle_kernel.values['EVEN_K'] = lambda args: True⋯ 1 unchanged linesexcept (AttributeError, KeyError, TypeError) as _e:print(f"[patch] heuristics failed: {_e}", file=_sys.stderr, flush=True)- # Set HIP_FORCE_DEV_KERNARG before any kernel launch_os.environ["HIP_FORCE_DEV_KERNARG"] = "1"- # --- HIP reduce kernel (replaces Triton reduce for KSPLIT>1 — lower launch overhead) ---+ # --- Modify _mxfp4_quant_op: replace log2/floor/exp2 with integer bit ops ---+ # The E8M0 scale computation uses tl.log2(amax).floor() - 2, but amax is already+ # a power of 2 (mantissa zeroed by & 0xFF800000). So the exponent can be extracted+ # with integer bit shifts, eliminating expensive v_log_f32 and v_ldexp_f32 SFU instructions.+ print("[fastquant] Modifying _mxfp4_quant_op source...", file=_sys.stderr, flush=True)+ try:+ _jit_fn = _gemm_a16wfp4_preshuffle_kernel.fn if hasattr(_gemm_a16wfp4_preshuffle_kernel, 'fn') else _gemm_a16wfp4_preshuffle_kernel+ _quant_fn = _jit_fn.__globals__['_mxfp4_quant_op']+ # _mxfp4_quant_op is a JITFunction — do NOT unwrap via .fn+ _old_qsrc = _quant_fn._src++ # Replacement 1: log2(amax).floor() → integer bit extraction+ _new_qsrc = _old_qsrc.replace(+ " amax = amax.to(tl.float32, bitcast=True)\n"+ " scale_e8m0_unbiased = tl.log2(amax).floor() - 2",+ " amax_exp = (amax >> 23) & 0xFF\n"+ " scale_e8m0_unbiased = (amax_exp.to(tl.int32) - 129).to(tl.float32)"+ )++ # Replacement 2: exp2(-scale) → integer FP32 construction+ _new_qsrc = _new_qsrc.replace(+ " quant_scale = tl.exp2(-scale_e8m0_unbiased)",+ " quant_scale = (((127.0 - scale_e8m0_unbiased).to(tl.int32).to(tl.uint32) << 23)).to(tl.float32, bitcast=True)"+ )++ if _new_qsrc != _old_qsrc:+ if hasattr(_quant_fn, '_unsafe_update_src'):+ _quant_fn._unsafe_update_src(_new_qsrc)+ else:+ _quant_fn._src = _new_qsrc+ if hasattr(_quant_fn, 'src'):+ _quant_fn.src = _new_qsrc+ if hasattr(_quant_fn, 'hash'):+ _quant_fn.hash = None+ # Also modify the KERNEL source to bust its Triton cache key+ # (quant is a dependency but its hash isn't in the kernel's cache key)+ _old_ksrc = _jit_fn._src+ _new_ksrc = _old_ksrc.replace(+ 'accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")',+ 'accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", acc=accumulator)'+ )+ if _new_ksrc != _old_ksrc:+ _jit_fn._unsafe_update_src(_new_ksrc)+ print("[fastquant] Applied quant + kernel source modifications", file=_sys.stderr, flush=True)+ else:+ print("[fastquant] Applied quant mod, kernel mod FAILED", file=_sys.stderr, flush=True)+ # Verify+ _vq = _quant_fn._src if hasattr(_quant_fn, '_src') else ''+ _vk = _jit_fn._src if hasattr(_jit_fn, '_src') else ''+ print(f"[fastquant] quant has amax_exp: {'amax_exp' in _vq}, kernel has acc=: {'acc=accumulator' in _vk}",+ file=_sys.stderr, flush=True)+ else:+ print("[fastquant] WARNING: replacement strings not found — source unchanged", file=_sys.stderr, flush=True)+ except Exception as _e:+ import traceback+ print(f"[fastquant] FAILED: {_e}", file=_sys.stderr, flush=True)+ traceback.print_exc(file=_sys.stderr)+ # --- End quant modification ---+++ # --- HIP reduce kernel ---_HIP_REDUCE_SRC = r"""#include <hip/hip_runtime.h>- // Manual bf16 conversion (round-to-nearest-even, matches Triton's .to(bf16))__device__ __forceinline__ unsigned short f32_to_bf16(float f) {unsigned int u;__builtin_memcpy(&u, &f, sizeof(u));⋯ 2 unchanged lines}template <int KSPLIT>- __global__ void reduce_k(const float* __restrict__ pp,- unsigned short* __restrict__ out, int MN) {- int idx = blockIdx.x * blockDim.x + threadIdx.x;- if (idx < MN) {- float s = pp[idx];+ __global__ void reduce_k_vec4(const float* __restrict__ pp,+ unsigned short* __restrict__ out, int MN) {+ int idx4 = (blockIdx.x * blockDim.x + threadIdx.x) * 4;+ if (idx4 + 3 < MN) {+ float4 s = *reinterpret_cast<const float4*>(pp + idx4);#pragma unroll- for (int k = 1; k < KSPLIT; k++) s += pp[k * MN + idx];- out[idx] = f32_to_bf16(s);+ for (int k = 1; k < KSPLIT; k++) {+ float4 v = *reinterpret_cast<const float4*>(pp + k * MN + idx4);+ s.x += v.x; s.y += v.y; s.z += v.z; s.w += v.w;+ }+ unsigned short r0 = f32_to_bf16(s.x);+ unsigned short r1 = f32_to_bf16(s.y);+ unsigned short r2 = f32_to_bf16(s.z);+ unsigned short r3 = f32_to_bf16(s.w);+ *reinterpret_cast<unsigned long long*>(out + idx4) =+ (unsigned long long)r0 | ((unsigned long long)r1 << 16) |+ ((unsigned long long)r2 << 32) | ((unsigned long long)r3 << 48);+ } else {+ for (int i = idx4; i < MN && i < idx4 + 4; i++) {+ float s = pp[i];+ #pragma unroll+ for (int k = 1; k < KSPLIT; k++) s += pp[k * MN + i];+ out[i] = f32_to_bf16(s);+ }}}⋯ 9 unchanged linesvoid reduce_op(torch::Tensor pp, torch::Tensor out, int M, int N, int ksplit) {int MN = M * N;- const int threads = 256;- const int blocks = (MN + threads - 1) / threads;const float* pp_ptr = pp.data_ptr<float>();unsigned short* out_ptr = reinterpret_cast<unsigned short*>(out.data_ptr());-+ const int threads_v = 64;+ const int elems_per_block = threads_v * 4;+ const int blocks_v = (MN + elems_per_block - 1) / elems_per_block;switch (ksplit) {- case 2: reduce_k<2><<<blocks, threads>>>(pp_ptr, out_ptr, MN); break;- case 3: reduce_k<3><<<blocks, threads>>>(pp_ptr, out_ptr, MN); break;- case 4: reduce_k<4><<<blocks, threads>>>(pp_ptr, out_ptr, MN); break;- case 7: reduce_k<7><<<blocks, threads>>>(pp_ptr, out_ptr, MN); break;- case 8: reduce_k<8><<<blocks, threads>>>(pp_ptr, out_ptr, MN); break;- default: reduce_k_gen<<<blocks, threads>>>(pp_ptr, out_ptr, MN, ksplit); break;+ case 2: reduce_k_vec4<2><<<blocks_v, threads_v>>>(pp_ptr, out_ptr, MN); break;+ case 3: reduce_k_vec4<3><<<blocks_v, threads_v>>>(pp_ptr, out_ptr, MN); break;+ case 4: reduce_k_vec4<4><<<blocks_v, threads_v>>>(pp_ptr, out_ptr, MN); break;+ case 7: reduce_k_vec4<7><<<blocks_v, threads_v>>>(pp_ptr, out_ptr, MN); break;+ case 8: reduce_k_vec4<8><<<blocks_v, threads_v>>>(pp_ptr, out_ptr, MN); break;+ default: {+ const int threads = 256;+ const int blocks = (MN + threads - 1) / threads;+ reduce_k_gen<<<blocks, threads>>>(pp_ptr, out_ptr, MN, ksplit);+ break;+ }}}"""⋯ 21 unchanged lines# --- End HIP reduce kernel ---- # --- Helper functions (needed before pre-warming) ---+ # --- Helper functions ---def _get_splitk(K: int, BLOCK_SIZE_K: int, NUM_KSPLIT: int):- """Adjust KSPLIT/BLOCK_K for EVEN_K alignment (inlined from aiter)."""SPLITK_BLOCK_SIZE = (triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K)⋯ 32 unchanged linesK = K_real // 2if M <= 32:- # Buckets 1-4: M≤32, BM=8, dynamic KSPLIT/BK- # B1: K=512 → KSPLIT=1, BK=256 (2 K-iters, pipeline)- # B2: K=1536 → KSPLIT=3, BK=256 (2 K-iters per split)- # B3: K=2048 → KSPLIT=4 or 2, BK=256 (1 or 2 K-iters per split)- # B4: K≥4096 → KSPLIT=7, BK=512 (1 K-iter per split)BLOCK_M = 8BLOCK_N = 128tiles_128 = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + 127) // 128)⋯ 1 unchanged linesif K_real >= 4096:KSPLIT = 7elif K_real >= 2048:- # Large-tile shapes: KSPLIT=2 BK=256 gives 2 K-iters (50% pipeline)- # vs KSPLIT=4 BK=256 with 1 K-iter. Less reduce (nk_pow2=2 vs 4).- # Only when BN=128 preserved (tiles*2 >= 3/4*CU) and wpe=1 (tiles*2 <= CU)if tiles_128 * 2 >= (_CU * 3) // 4 and tiles_128 * 2 <= _CU:KSPLIT = 2else:KSPLIT = 4elif K_real >= 1536:- # Same logic: KSPLIT=2 gives 2 K-iters vs KSPLIT=3 with 1 K-iterif tiles_128 * 2 >= (_CU * 3) // 4 and tiles_128 * 2 <= _CU:KSPLIT = 2else:⋯ 9 unchanged lines"cache_modifier": ".cg", "NUM_KSPLIT": KSPLIT,}else:- # Buckets 5-8: M>32- # B5: M=64 low CU util → BM=8, dynamic KSPLIT- # B6: M=64 high CU util → BM=16, KSPLIT=1-2- # B7: M=128 → BM=8 or 16, KSPLIT=1-2- # B8: M=256 → BM=16, KSPLIT=1BLOCK_M = 16if M <= 128:tiles_bm16 = ((M + 15) // 16) * ((N + 127) // 128)⋯ 56 unchanged linesif cfg["NUM_KSPLIT"] > 1:grid_reduce = (triton.cdiv(M, 16), triton.cdiv(N, 16))- result = (cfg, actual_ksplit, nk_pow2, grid_main, grid_reduce)+ result = (cfg, actual_ksplit, nk_pow2, grid_main, grid_reduce,+ K, cfg["BLOCK_SIZE_M"], cfg["BLOCK_SIZE_N"], cfg["BLOCK_SIZE_K"],+ cfg["NUM_KSPLIT"], cfg["SPLITK_BLOCK_SIZE"], cfg["waves_per_eu"])_CFG_CACHE[key] = resultreturn result- # --- Nuclear pre-warming framework ---- # Enumerate ALL unique Triton cache keys across 171 shapes.- # Phase 1: compile K>=1536 M≤32 configs WITHOUT disable-lsr (these regress +1.8% with it).- # Phase 2: set DISABLE_LLVM_OPT=disable-lsr (helps M>32 shapes -2.5%).- # Phase 3: compile remaining configs WITH disable-lsr (with 200s timeout safety).- # Phase 4: compile reduce kernel configs.+ # --- Nuclear pre-warming ---_WARMUP_T0 = _time.time()_PREWARMED_CONFIGS = {}- # Collect unique cache keys- _NO_LSR = {} # M≤32 K>=1536 → compile without disable-lsr- _LSR = {} # everything else → compile with disable-lsr- _REDUCE = set() # (actual_ksplit, nk_pow2) for reduce kernel+ _NO_LSR = {}+ _LSR = {}+ _REDUCE = set()for _nw, _kw in _NK_FAMILIES:for _mw in _M_VALUES:- _cw, _aw, _nkw, _, _ = _get_cfg(_mw, _nw, _kw)+ _cw, _aw, _nkw, _, _, _, _, _, _, _, _, _ = _get_cfg(_mw, _nw, _kw)_ck = (_cw["BLOCK_SIZE_M"], _cw["BLOCK_SIZE_N"], _cw["BLOCK_SIZE_K"],_cw["NUM_KSPLIT"], _cw["SPLITK_BLOCK_SIZE"], _cw["waves_per_eu"])if _mw <= 32 and _kw >= 1536:⋯ 3 unchanged linesif _aw is not None:_REDUCE.add((_aw, _nkw))- # Configs in both groups: keep in no-lsr (K=7168 M≤32 needs no-lsr)for _k in _NO_LSR:_LSR.pop(_k, None)print(f"[pre-warm] {len(_NO_LSR)} no-lsr + {len(_LSR)} lsr GEMM, {len(_REDUCE)} reduce configs",file=_sys.stderr, flush=True)- # Dummy tensors (oversized to avoid OOB on any config)_wA = torch.zeros(32, 8192, dtype=torch.bfloat16, device="cuda")_wBw = torch.zeros(16, 65536, dtype=torch.uint8, device="cuda")_wBs = torch.zeros(16, 65536, dtype=torch.uint8, device="cuda")⋯ 2 unchanged linesdef _pw(bm, bn, bk, ks, spk, wpe):- """Pre-warm one GEMM config by launching with dummy data."""c = {"BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": bn, "BLOCK_SIZE_K": bk,"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,"waves_per_eu": wpe, "matrix_instr_nonkdim": 16,⋯ 24 unchanged linesprint(f"[pre-warm] Phase 2: DISABLE_LLVM_OPT=disable-lsr set ({_time.time()-_WARMUP_T0:.0f}s)",file=_sys.stderr, flush=True)- # Phase 3: remaining GEMM configs with disable-lsr (timeout safety: 200s total)+ # Phase 3: remaining GEMM configs with disable-lsr_lsr_list = sorted(_LSR)print(f"[pre-warm] Phase 3: {len(_lsr_list)} remaining GEMM configs (disable-lsr)...",file=_sys.stderr, flush=True)⋯ 32 unchanged linesprint(f"[pre-warm] Done: {len(_PREWARMED_CONFIGS)} GEMM configs in {_time.time()-_WARMUP_T0:.0f}s",file=_sys.stderr, flush=True)- _gc.disable() # Prevent GC pauses during benchmark+ _gc.disable()# --- End pre-warming ---⋯ 21 unchanged linesif not A.is_contiguous():A = A.contiguous()- shape_prefix = tuple(A.shape[:-1])- A_2d = A.view(-1, A.shape[-1])- M = A_2d.shape[0]+ _ndim = A.ndim+ if _ndim == 2:+ A_2d = A+ M = A.shape[0]+ else:+ A_2d = A.view(-1, A.shape[-1])+ M = A_2d.shape[0]N = data[3].shape[0]K_bytes = data[3].shape[1]K_real = K_bytes * 2- K = K_real // 2- cfg, actual_ksplit, nk_pow2, grid_main, grid_reduce = _get_cfg(M, N, K_real)+ cfg, actual_ksplit, nk_pow2, grid_main, grid_reduce, K, BM, BN, BK, KS, SPK, WPE = _get_cfg(M, N, K_real)- # Per-shape logging (first call only)_sk = (M, N, K_real)if _sk not in _LOGGED:_LOGGED.add(_sk)- print(f"[kernel] M={M} N={N} K={K_real} BM={cfg['BLOCK_SIZE_M']} BN={cfg['BLOCK_SIZE_N']} "- f"BK={cfg['BLOCK_SIZE_K']} KS={cfg['NUM_KSPLIT']} wpe={cfg['waves_per_eu']} "+ print(f"[kernel] M={M} N={N} K={K_real} BM={BM} BN={BN} "+ f"BK={BK} KS={KS} wpe={WPE} "f"grid={grid_main[0]}", file=_sys.stderr, flush=True)- dev = A.device- okey = (dev.index, M, N)+ okey = (M, N)if okey not in _OUT_BUF:- _OUT_BUF[okey] = torch.empty((M, N), dtype=torch.bfloat16, device=dev)+ _OUT_BUF[okey] = torch.empty((M, N), dtype=torch.bfloat16, device="cuda")y = _OUT_BUF[okey]B_w, B_s, stride_bw0, stride_bs0 = _get_preshuffle_b(data)- if cfg["NUM_KSPLIT"] > 1:- ppkey = (dev.index, nk_pow2, M, N)+ if KS > 1:+ ppkey = (nk_pow2, M, N)if ppkey not in _YPP_BUF:_YPP_BUF[ppkey] = torch.empty(- (nk_pow2, M, N), dtype=torch.float32, device=dev+ (nk_pow2, M, N), dtype=torch.float32, device="cuda")y_pp = _YPP_BUF[ppkey]stride_ck = M * N⋯ 12 unchanged linesstride_bw0, 1,stride_ck, stride_cm, 1,stride_bs0, 1,+ BLOCK_SIZE_M=BM, BLOCK_SIZE_N=BN, BLOCK_SIZE_K=BK,+ GROUP_SIZE_M=1, NUM_KSPLIT=KS, SPLITK_BLOCK_SIZE=SPK,+ num_warps=4, num_stages=2, waves_per_eu=WPE,+ matrix_instr_nonkdim=16, cache_modifier=".cg",PREQUANT=True,- **cfg,)if y_pp is not None:⋯ 8 unchanged linesactual_ksplit, nk_pow2,)- return y.view(*shape_prefix, N)+ if _ndim == 2:+ return y+ return y.view(*A.shape[:-1], N)
scrolls · 391 diff lines total
Best evidence level for this revision: reported
JSON