submission 754371
.jonnss · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 537 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-754371?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:a8877a15a922f0ce34cbcec5af4546266b27407cbdb39a4c62de2ac2d86b1632
license declaredunknown
license concludedunknown
authors.jonnss
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
CDNA4 blocked FP4 matmul via aiter preshuffle path.num-warps = 4
num_warps=4, num_stages=2, waves_per_eu=wpe,split-k
_rows = ["cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio"]stages = 2
num_warps=4, num_stages=2, waves_per_eu=wpe,vector-width = float4
float4 acc = *reinterpret_cast<const float4*>(src + base);Kernel source
submission.py537 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
CDNA4 blocked FP4 matmul via aiter preshuffle path.
Runtime patches the quantization subroutine with ISA-level
paired conversion and adjusts the accumulator idiom.
Occupancy-aware tiling with phased JIT warmup.
"""
import os as _env
_env.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
_env.environ.setdefault("CXX", "clang++")
import uuid as _uid
_env.environ["TRITON_CACHE_DIR"] = f"/tmp/_tc_{_uid.uuid4().hex[:8]}"
# Synthesize a minimal config CSV so aiter skips its heavy build paths
_ASM_LABEL = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_CSV_TMP = "/tmp/_fp4_cfg.csv"
_ENGINES = 256
_DIM_PAIRS = [
(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),
]
_BATCH_DIMS = [1, 2, 4, 8, 16, 32, 64, 128, 256]
_rows = ["cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio"]
for _d1, _d2 in _DIM_PAIRS:
for _b in _BATCH_DIMS:
_wave_cnt = ((_b + 31) // 32) * ((_d1 + 127) // 128)
_ratio = _ENGINES / max(_wave_cnt, 1)
_lg = 0
while _ratio >= pow(2, _lg + 1) and (pow(2, _lg + 1) * 128) < 2 * _d2:
_lg += 1
_lg = min(_lg, 3)
_rows.append(f"{_ENGINES},{_b},{_d1},{_d2},21,{_lg},1.0,{_ASM_LABEL},0,0,0.0")
with open(_CSV_TMP, "w") as _fh:
_fh.write("\n".join(_rows))
_env.environ["AITER_CONFIG_GEMM_A4W4"] = (
_CSV_TMP + ":/home/runner/aiter/aiter/configs/a4w4_blockscale_tuned_gemm.csv"
)
import torch
torch.set_grad_enabled(False)
import triton
import triton.language as tl
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 _io
import time as _clk
import gc as _mem
_io.setswitchinterval(1.0)
_log = lambda s: print(s, file=_io.stderr, flush=True)
# ---- Override heuristics to avoid dynamic tile selection ----
try:
_gemm_a16wfp4_preshuffle_kernel.values['GRID_MN'] = lambda args: 1
_gemm_a16wfp4_preshuffle_kernel.values['EVEN_K'] = lambda args: True
_log("[setup] heuristics locked")
except Exception as _exc:
_log(f"[setup] heuristic override failed: {_exc}")
_env.environ["HIP_FORCE_DEV_KERNARG"] = "1"
# ---- Inject ISA-level BF16->FP4 quantization ----
_log("[setup] patching quantizer with hardware conversion...")
try:
_kern_obj = (
_gemm_a16wfp4_preshuffle_kernel.fn
if hasattr(_gemm_a16wfp4_preshuffle_kernel, 'fn')
else _gemm_a16wfp4_preshuffle_kernel
)
_quant_ref = _kern_obj.__globals__['_mxfp4_quant_op']
_patched_body = '''def _mxfp4_quant_op(
x,
BLOCK_SIZE_N,
BLOCK_SIZE_M,
MXFP4_QUANT_BLOCK_SIZE,
):
"""ISA-accelerated BF16 to packed FP4 via v_cvt_scalef32_pk_fp4_bf16."""
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
HALF_BLOCK: tl.constexpr = MXFP4_QUANT_BLOCK_SIZE // 2
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_exp = (amax >> 23) & 0xFF
scale_e8m0_unbiased = (amax_exp.to(tl.int32) - 129).to(tl.float32)
scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
bs_e8m0 = (scale_e8m0_unbiased + 127).to(tl.float32).to(tl.uint8)
biased_exp_f = tl.maximum(scale_e8m0_unbiased + 127.0, 1.0)
hw_scale = (biased_exp_f.to(tl.int32).to(tl.uint32) << 23).to(tl.float32, bitcast=True)
x_bf16 = x.to(tl.bfloat16)
x_pairs = x_bf16.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, HALF_BLOCK, 2)
evens, odds = tl.split(x_pairs)
lo = evens.to(tl.uint16, bitcast=True).to(tl.uint32)
hi = odds.to(tl.uint16, bitcast=True).to(tl.uint32)
packed_bf16 = lo | (hi << 16)
result = tl.inline_asm_elementwise(
"v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2",
"=v,v,v",
[packed_bf16, hw_scale],
dtype=tl.uint32,
is_pure=True,
pack=1,
)
x_fp4 = (result & 0xFF).to(tl.uint8)
x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
'''
if hasattr(_quant_ref, '_unsafe_update_src'):
_quant_ref._unsafe_update_src(_patched_body)
else:
_quant_ref._src = _patched_body
if hasattr(_quant_ref, 'src'):
_quant_ref.src = _patched_body
if hasattr(_quant_ref, 'hash'):
_quant_ref.hash = None
_orig_src = _kern_obj._src
_mod_src = _orig_src.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 _mod_src != _orig_src:
_kern_obj._unsafe_update_src(_mod_src)
_log("[setup] quantizer + kernel patched OK")
else:
_log("[setup] quantizer patched, kernel string mismatch")
_check = _quant_ref._src if hasattr(_quant_ref, '_src') else ''
_log(f"[setup] has hw asm: {'inline_asm_elementwise' in _check}")
except Exception as _exc:
import traceback
_log(f"[setup] patch FAILED: {_exc}")
traceback.print_exc(file=_io.stderr)
# ---- Native HIP accumulator merger for K-partitioned runs ----
_MERGER_HIP = r"""
#include <hip/hip_runtime.h>
__device__ __forceinline__ unsigned short to_bf16(float val) {
unsigned int raw;
__builtin_memcpy(&raw, &val, sizeof(raw));
unsigned int bias = ((raw >> 16) & 1) + 0x7FFFu;
return (unsigned short)((raw + bias) >> 16);
}
template <int NP>
__global__ void vec4_sum(const float* __restrict__ src,
unsigned short* __restrict__ dst, int len) {
int base = (blockIdx.x * blockDim.x + threadIdx.x) * 4;
if (base + 3 < len) {
float4 acc = *reinterpret_cast<const float4*>(src + base);
#pragma unroll
for (int p = 1; p < NP; p++) {
float4 part = *reinterpret_cast<const float4*>(src + p * len + base);
acc.x += part.x; acc.y += part.y; acc.z += part.z; acc.w += part.w;
}
unsigned short a = to_bf16(acc.x), b = to_bf16(acc.y);
unsigned short c = to_bf16(acc.z), d = to_bf16(acc.w);
*reinterpret_cast<unsigned long long*>(dst + base) =
(unsigned long long)a | ((unsigned long long)b << 16) |
((unsigned long long)c << 32) | ((unsigned long long)d << 48);
} else {
for (int j = base; j < len && j < base + 4; j++) {
float acc = src[j];
#pragma unroll
for (int p = 1; p < NP; p++) acc += src[p * len + j];
dst[j] = to_bf16(acc);
}
}
}
__global__ void scalar_sum(const float* __restrict__ src,
unsigned short* __restrict__ dst,
int len, int np) {
int gid = blockIdx.x * blockDim.x + threadIdx.x;
if (gid < len) {
float acc = src[gid];
for (int p = 1; p < np; p++) acc += src[p * len + gid];
dst[gid] = to_bf16(acc);
}
}
void merge_partials(torch::Tensor src, torch::Tensor dst, int R, int C, int np) {
int len = R * C;
const float* sp = src.data_ptr<float>();
unsigned short* dp = reinterpret_cast<unsigned short*>(dst.data_ptr());
const int thr = 64, stride = thr * 4;
const int nblk = (len + stride - 1) / stride;
switch (np) {
case 2: vec4_sum<2><<<nblk, thr>>>(sp, dp, len); break;
case 3: vec4_sum<3><<<nblk, thr>>>(sp, dp, len); break;
case 4: vec4_sum<4><<<nblk, thr>>>(sp, dp, len); break;
case 7: vec4_sum<7><<<nblk, thr>>>(sp, dp, len); break;
case 8: vec4_sum<8><<<nblk, thr>>>(sp, dp, len); break;
default: {
const int t2 = 256, b2 = (len + t2 - 1) / t2;
scalar_sum<<<b2, t2>>>(sp, dp, len, np);
break;
}
}
}
"""
_MERGER_HDR = "void merge_partials(torch::Tensor src, torch::Tensor dst, int R, int C, int np);"
_HAS_HIP_MERGER = False
try:
from torch.utils.cpp_extension import load_inline as _jit
_jit_t0 = _clk.time()
_hip_merger = _jit(
name="fp4_kmerge",
cpp_sources=[_MERGER_HDR],
cuda_sources=[_MERGER_HIP],
functions=["merge_partials"],
verbose=False,
extra_cuda_cflags=["--offload-arch=gfx950", "-O3"],
)
_HAS_HIP_MERGER = True
_log(f"[setup] HIP merger ready ({_clk.time()-_jit_t0:.1f}s)")
except Exception as _exc:
_log(f"[setup] HIP merger unavailable: {_exc}")
# ---- Partition alignment utility ----
def _align_parts(kh, bk, np):
span = triton.cdiv((2 * triton.cdiv(kh, np)), bk) * bk
while np > 1 and bk > 16:
ok = (kh % (span // 2) == 0 and span % bk == 0 and kh % (bk // 2) == 0)
if ok:
break
elif kh % (span // 2) != 0 and np > 1:
np //= 2
elif span % bk != 0:
np = np // 2 if np > 1 else np
if np <= 1 and bk > 16:
bk //= 2
elif kh % (bk // 2) != 0 and bk > 16:
bk //= 2
else:
break
span = triton.cdiv((2 * triton.cdiv(kh, np)), bk) * bk
return span, bk, np
# ---- Occupancy-driven config resolver ----
_resolved = {}
def _shape_config(batch, cols, depth):
tag = (batch, cols, depth)
if tag in _resolved:
return _resolved[tag]
kh = depth // 2
if batch <= 32:
bm, bn = 8, 128
wave_est = ((batch + bm - 1) // bm) * ((cols + 127) // 128)
np = 1
if depth >= 4096:
np = 7
elif depth >= 2048:
np = 2 if (wave_est * 2 >= (_ENGINES * 3) // 4 and wave_est * 2 <= _ENGINES) else 4
elif depth >= 1536:
np = 2 if (wave_est * 2 >= (_ENGINES * 3) // 4 and wave_est * 2 <= _ENGINES) else 3
bk = 256 if depth <= np * 512 or (np == 2 and depth <= np * 1024) else 512
if wave_est * np < (_ENGINES * 3) // 4:
bn = 64
total_wg = ((batch + bm - 1) // bm) * ((cols + bn - 1) // bn) * np
wpe = 2 if total_wg > _ENGINES else 1
else:
bm = 16
if batch <= 128:
est16 = ((batch + 15) // 16) * ((cols + 127) // 128)
if est16 < (_ENGINES * 3) // 4:
bm = 8
wave_est = ((batch + bm - 1) // bm) * ((cols + 127) // 128)
bn, np = 128, 1
if _ENGINES // 2 <= wave_est <= _ENGINES and (depth >= 7168 or (depth >= 2048 and bm == 8)):
np = 2
elif wave_est < _ENGINES // 2 and depth > 512:
if depth >= 4096:
np = 2 if wave_est * 2 >= _ENGINES else 7
elif depth >= 2048:
np = 2
elif depth >= 1536:
np = 3
bk = 256 if depth <= max(np * 4096, 2048) else 512
if wave_est * np < (_ENGINES * 3) // 4:
bn = 64
total_wg = ((batch + bm - 1) // bm) * ((cols + bn - 1) // bn) * np
wpe = 2 if total_wg > _ENGINES else 1
params = {
"BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": max(bn, 32), "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": np,
}
if params["NUM_KSPLIT"] > 1:
span, bk2, np2 = _align_parts(kh, params["BLOCK_SIZE_K"], params["NUM_KSPLIT"])
params["SPLITK_BLOCK_SIZE"] = span
params["BLOCK_SIZE_K"] = bk2
params["NUM_KSPLIT"] = np2
if params["BLOCK_SIZE_K"] >= 2 * kh:
params["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * kh)
params["SPLITK_BLOCK_SIZE"] = 2 * kh
params["NUM_KSPLIT"] = 1
params["BLOCK_SIZE_N"] = max(params["BLOCK_SIZE_N"], 32)
if params["NUM_KSPLIT"] == 1:
params["SPLITK_BLOCK_SIZE"] = 2 * kh
real_np, padded_np = None, None
if params["NUM_KSPLIT"] > 1:
real_np = triton.cdiv(kh, params["SPLITK_BLOCK_SIZE"] // 2)
padded_np = triton.next_power_of_2(params["NUM_KSPLIT"])
m_tiles = triton.cdiv(batch, params["BLOCK_SIZE_M"])
n_tiles = triton.cdiv(cols, params["BLOCK_SIZE_N"])
launch_grid = (params["NUM_KSPLIT"] * m_tiles * n_tiles,)
red_grid = None
if params["NUM_KSPLIT"] > 1:
red_grid = (triton.cdiv(batch, 16), triton.cdiv(cols, 16))
bundle = (
params, real_np, padded_np, launch_grid, red_grid,
kh, params["BLOCK_SIZE_M"], params["BLOCK_SIZE_N"],
params["BLOCK_SIZE_K"], params["NUM_KSPLIT"],
params["SPLITK_BLOCK_SIZE"], params["waves_per_eu"],
)
_resolved[tag] = bundle
return bundle
# ---- Phased JIT warmup ----
_warm_t0 = _clk.time()
_warmed = {}
_phase1 = {}
_phase2 = {}
_red_set = set()
for _dn, _dk in _DIM_PAIRS:
for _db in _BATCH_DIMS:
_cb, _rn, _rp, _, _, _, _, _, _, _, _, _ = _shape_config(_db, _dn, _dk)
_sig = (
_cb["BLOCK_SIZE_M"], _cb["BLOCK_SIZE_N"], _cb["BLOCK_SIZE_K"],
_cb["NUM_KSPLIT"], _cb["SPLITK_BLOCK_SIZE"], _cb["waves_per_eu"],
)
if _db <= 32 and _dk >= 1536:
_phase1.setdefault(_sig, True)
else:
_phase2.setdefault(_sig, True)
if _rn is not None:
_red_set.add((_rn, _rp))
for _s in _phase1:
_phase2.pop(_s, None)
_log(f"[warm] {len(_phase1)} phase1 + {len(_phase2)} phase2, {len(_red_set)} reducers")
_dummy_x = torch.zeros(32, 8192, dtype=torch.bfloat16, device="cuda")
_dummy_w = torch.zeros(16, 65536, dtype=torch.uint8, device="cuda")
_dummy_s = torch.zeros(16, 65536, dtype=torch.uint8, device="cuda")
_dummy_pp = torch.zeros(16, 32, 256, dtype=torch.float32, device="cuda")
_dummy_y = torch.zeros(32, 256, dtype=torch.bfloat16, device="cuda")
def _fire_warmup(bm, bn, bk, ks, spk, wpe):
cfg = {
"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,
}
target = _dummy_pp if ks > 1 else _dummy_y
_gemm_a16wfp4_preshuffle_kernel[(max(ks, 1),)](
_dummy_x, _dummy_w, target, _dummy_s, bm, bn, spk // 2,
_dummy_x.stride(0), _dummy_x.stride(1),
_dummy_w.stride(0), _dummy_w.stride(1),
0 if ks <= 1 else _dummy_pp.stride(0),
_dummy_y.stride(0) if ks <= 1 else _dummy_pp.stride(1),
_dummy_y.stride(1) if ks <= 1 else _dummy_pp.stride(2),
_dummy_s.stride(0), _dummy_s.stride(1),
PREQUANT=True, **cfg,
)
_log("[warm] phase 1 (no lsr)...")
for _sig in sorted(_phase1):
try:
_fire_warmup(*_sig)
_warmed[_sig] = 1
_log(f" {_sig[0]}x{_sig[1]}x{_sig[2]} ks={_sig[3]} ({_clk.time()-_warm_t0:.0f}s)")
except Exception as _exc:
_log(f" {_sig}: ERR {_exc}")
_env.environ["DISABLE_LLVM_OPT"] = "disable-lsr"
_log(f"[warm] phase 2 (with lsr) @ {_clk.time()-_warm_t0:.0f}s")
for _idx, _sig in enumerate(sorted(_phase2)):
if _clk.time() - _warm_t0 > 200:
_log(f" timeout, {len(_phase2) - _idx} skipped")
break
try:
_fire_warmup(*_sig)
_warmed[_sig] = 2
_log(f" {_sig[0]}x{_sig[1]}x{_sig[2]} ks={_sig[3]} ({_clk.time()-_warm_t0:.0f}s)")
except Exception as _exc:
_log(f" {_sig}: ERR {_exc}")
_log(f"[warm] reducers...")
for _rn, _rp in sorted(_red_set):
if _clk.time() - _warm_t0 > 230:
_log(" timeout")
break
try:
_gemm_afp4wfp4_reduce_kernel[(1, 1)](
_dummy_pp, _dummy_y, 16, 16,
_dummy_pp.stride(0), _dummy_pp.stride(1), _dummy_pp.stride(2),
_dummy_y.stride(0), _dummy_y.stride(1), 16, 16, _rn, _rp,
)
except Exception:
pass
del _dummy_x, _dummy_w, _dummy_s, _dummy_pp, _dummy_y, _fire_warmup
del _phase1, _phase2, _red_set
torch.cuda.empty_cache()
_log(f"[warm] done: {len(_warmed)} configs in {_clk.time()-_warm_t0:.0f}s")
_mem.disable()
# ---- Runtime dispatch state ----
_wt_cache = {}
_dest_buf = {}
_frag_buf = {}
_seen = set()
def _prepare_wt(data):
addr = data[3].data_ptr()
if addr not in _wt_cache:
n_dim = data[3].shape[0]
k_bytes = data[3].shape[1]
sr, sc = data[4].shape
n_grp = n_dim // 32
w_view = data[3].view(torch.uint8).reshape(n_dim // 16, k_bytes * 16)
s_view = data[4].view(torch.uint8).reshape(sr // 32, sc * 32)[:n_grp].contiguous()
_wt_cache[addr] = (w_view, s_view, w_view.stride(0), s_view.stride(0))
return _wt_cache[addr]
def custom_kernel(data: input_t) -> output_t:
X = data[0]
if not X.is_contiguous():
X = X.contiguous()
ndims = X.ndim
X_flat = X if ndims == 2 else X.view(-1, X.shape[-1])
batch = X_flat.shape[0]
cols = data[3].shape[0]
depth = data[3].shape[1] * 2
(params, real_np, padded_np, launch_grid, red_grid,
kh, bm, bn, bk, ks, spk, wpe) = _shape_config(batch, cols, depth)
tag = (batch, cols, depth)
if tag not in _seen:
_seen.add(tag)
_log(f"[run] {batch}x{cols}x{depth} bm={bm} bn={bn} bk={bk} ks={ks} wpe={wpe}")
okey = (batch, cols)
if okey not in _dest_buf:
_dest_buf[okey] = torch.empty((batch, cols), dtype=torch.bfloat16, device="cuda")
dest = _dest_buf[okey]
w_view, s_view, sw0, ss0 = _prepare_wt(data)
if ks > 1:
fkey = (padded_np, batch, cols)
if fkey not in _frag_buf:
_frag_buf[fkey] = torch.empty(
(padded_np, batch, cols), dtype=torch.float32, device="cuda"
)
frags = _frag_buf[fkey]
stride_p, stride_r = batch * cols, cols
else:
frags = None
stride_p, stride_r = 0, cols
_gemm_a16wfp4_preshuffle_kernel[launch_grid](
X_flat, w_view,
dest if frags is None else frags,
s_view, batch, cols, kh,
depth, 1, sw0, 1,
stride_p, stride_r, 1,
ss0, 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 frags is not None:
if _HAS_HIP_MERGER:
_hip_merger.merge_partials(frags, dest, batch, cols, real_np)
else:
_gemm_afp4wfp4_reduce_kernel[red_grid](
frags, dest, batch, cols,
batch * cols, cols, 1, cols, 1,
16, 16, real_np, padded_np,
)
return dest if ndims == 2 else dest.view(*X.shape[:-1], cols)
scrolls · 537 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 738857.
- """- Optimized MXFP4 GEMM submission based on GEMM-Reference.md best practices.+ #!POPCORN leaderboard amd-mxfp4-mm+ #!POPCORN gpu MI355X- Key optimizations (all proven via 200+ experiments):- 1. Integer E8M0 scale (replaces log2/floor/exp2 SFU instructions)- 2. .wt store + fast_math + acc=accumulator source patches- 3. Selective disable-lsr, GRID_MN/EVEN_K heuristic patches- 4. Nuclear pre-warming, wave scheduling, eviction_policy- 5. Tuned per-shape configs (BK=256 pipeline, KSPLIT routing, BM=8 for M<=32)"""- import gc- import importlib- import os- import re- import sys- import weakref+ CDNA4 blocked FP4 matmul via aiter preshuffle path.+ Runtime patches the quantization subroutine with ISA-level+ paired conversion and adjusts the accumulator idiom.+ Occupancy-aware tiling with phased JIT warmup.+ """- # ── Environment setup (before any imports that trigger Triton/HIP) ──────────- os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")- os.environ.setdefault("TRITON_HIP_ENABLE_WAVE_SCHEDULING", "1")+ import os as _env+ _env.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")+ _env.environ.setdefault("CXX", "clang++")+ import uuid as _uid+ _env.environ["TRITON_CACHE_DIR"] = f"/tmp/_tc_{_uid.uuid4().hex[:8]}"- import aiter+ # Synthesize a minimal config CSV so aiter skips its heavy build paths+ _ASM_LABEL = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"+ _CSV_TMP = "/tmp/_fp4_cfg.csv"+ _ENGINES = 256+ _DIM_PAIRS = [+ (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),+ ]+ _BATCH_DIMS = [1, 2, 4, 8, 16, 32, 64, 128, 256]+ _rows = ["cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio"]+ for _d1, _d2 in _DIM_PAIRS:+ for _b in _BATCH_DIMS:+ _wave_cnt = ((_b + 31) // 32) * ((_d1 + 127) // 128)+ _ratio = _ENGINES / max(_wave_cnt, 1)+ _lg = 0+ while _ratio >= pow(2, _lg + 1) and (pow(2, _lg + 1) * 128) < 2 * _d2:+ _lg += 1+ _lg = min(_lg, 3)+ _rows.append(f"{_ENGINES},{_b},{_d1},{_d2},21,{_lg},1.0,{_ASM_LABEL},0,0,0.0")+ with open(_CSV_TMP, "w") as _fh:+ _fh.write("\n".join(_rows))+ _env.environ["AITER_CONFIG_GEMM_A4W4"] = (+ _CSV_TMP + ":/home/runner/aiter/aiter/configs/a4w4_blockscale_tuned_gemm.csv"+ )+import torch+ torch.set_grad_enabled(False)import tritonimport triton.language as tl- from aiter import dtypes- from aiter.ops.triton.quant import dynamic_mxfp4_quant- from aiter.utility.fp4_utils import e8m0_shuffle-+ 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 _io+ import time as _clk+ import gc as _mem+ _io.setswitchinterval(1.0)+ _log = lambda s: print(s, file=_io.stderr, flush=True)- # ── Global state ────────────────────────────────────────────────────────────- _CU = 256- _LOW_UTIL_THRESHOLD = (_CU * 3) // 4- _PRESHUFFLE_CACHE = {}- _OUT_CACHE = {}- _PARTIAL_CACHE = {}- _SHAPE_CFG_CACHE = {}- _LOGGED_PATHS = set()+ # ---- Override heuristics to avoid dynamic tile selection ----+ try:+ _gemm_a16wfp4_preshuffle_kernel.values['GRID_MN'] = lambda args: 1+ _gemm_a16wfp4_preshuffle_kernel.values['EVEN_K'] = lambda args: True+ _log("[setup] heuristics locked")+ except Exception as _exc:+ _log(f"[setup] heuristic override failed: {_exc}")- _DIRECT_KERNEL = None- _REDUCE_KERNEL = None- _GET_SPLITK = None- _INIT_DONE = False- _QUANT_PATCHED = False- _KERNEL_PATCHED = False+ _env.environ["HIP_FORCE_DEV_KERNARG"] = "1"- gc.disable()- torch.set_grad_enabled(False)- sys.setswitchinterval(1.0)+ # ---- Inject ISA-level BF16->FP4 quantization ----+ _log("[setup] patching quantizer with hardware conversion...")+ try:+ _kern_obj = (+ _gemm_a16wfp4_preshuffle_kernel.fn+ if hasattr(_gemm_a16wfp4_preshuffle_kernel, 'fn')+ else _gemm_a16wfp4_preshuffle_kernel+ )+ _quant_ref = _kern_obj.__globals__['_mxfp4_quant_op']- # ── Helpers ─────────────────────────────────────────────────────────────────- def _ceil_div(a: int, b: int) -> int:- return (a + b - 1) // b+ _patched_body = '''def _mxfp4_quant_op(+ x,+ BLOCK_SIZE_N,+ BLOCK_SIZE_M,+ MXFP4_QUANT_BLOCK_SIZE,+ ):+ """ISA-accelerated BF16 to packed FP4 via v_cvt_scalef32_pk_fp4_bf16."""+ NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE+ HALF_BLOCK: tl.constexpr = MXFP4_QUANT_BLOCK_SIZE // 2+ x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)- def _view_dtype(tensor: torch.Tensor, dtype) -> torch.Tensor:- if tensor.dtype == dtype:- return tensor- return tensor.view(dtype)+ 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_exp = (amax >> 23) & 0xFF+ scale_e8m0_unbiased = (amax_exp.to(tl.int32) - 129).to(tl.float32)+ scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)- # ── Source patching ─────────────────────────────────────────────────────────+ bs_e8m0 = (scale_e8m0_unbiased + 127).to(tl.float32).to(tl.uint8)- def _patch_quant_op():- """- Replace _mxfp4_quant_op with integer E8M0 scale computation.- Eliminates log2/floor/exp2 SFU instructions (~24 instructions -> ~6).- """- global _QUANT_PATCHED- if _QUANT_PATCHED:- return+ biased_exp_f = tl.maximum(scale_e8m0_unbiased + 127.0, 1.0)+ hw_scale = (biased_exp_f.to(tl.int32).to(tl.uint32) << 23).to(tl.float32, bitcast=True)- try:- quant_mod = importlib.import_module("aiter.ops.triton.quant")- quant_fn = getattr(quant_mod, "_mxfp4_quant_op", None)- if quant_fn is None:- return+ x_bf16 = x.to(tl.bfloat16)+ x_pairs = x_bf16.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, HALF_BLOCK, 2)+ evens, odds = tl.split(x_pairs)+ lo = evens.to(tl.uint16, bitcast=True).to(tl.uint32)+ hi = odds.to(tl.uint16, bitcast=True).to(tl.uint32)+ packed_bf16 = lo | (hi << 16)- if not hasattr(quant_fn, 'src'):- quant_fn = _get_jit_fn(quant_fn)- if not hasattr(quant_fn, 'src'):- print(f"[mm-opt] Quant fn has no .src, type: {type(quant_fn).__name__}", file=sys.stderr, flush=True)- return+ result = tl.inline_asm_elementwise(+ "v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2",+ "=v,v,v",+ [packed_bf16, hw_scale],+ dtype=tl.uint32,+ is_pure=True,+ pack=1,+ )- src = quant_fn.src+ x_fp4 = (result & 0xFF).to(tl.uint8)+ x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)- pass # src loaded- new_src = src+ return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)+ '''- # 1. Replace log2(amax).floor() - 2 with integer bit extraction- # amax is already power of 2 (mantissa zeroed), so log2 = exponent - 127- new_src = new_src.replace(- "scale_e8m0_unbiased = tl.log2(amax).floor() - 2",- "amax_bits_int = amax.to(tl.int32, bitcast=True)\n"- " scale_e8m0_unbiased = (((amax_bits_int >> 23) & 0xFF).to(tl.float32) - 127.0 - 2)"- )+ if hasattr(_quant_ref, '_unsafe_update_src'):+ _quant_ref._unsafe_update_src(_patched_body)+ else:+ _quant_ref._src = _patched_body+ if hasattr(_quant_ref, 'src'):+ _quant_ref.src = _patched_body+ if hasattr(_quant_ref, 'hash'):+ _quant_ref.hash = None- # 2. Replace exp2 with integer bitcast- new_src = new_src.replace(- "quant_scale = tl.exp2(-scale_e8m0_unbiased)",- "neg_scale_int = (-scale_e8m0_unbiased).to(tl.int32)\n"- " quant_scale = ((neg_scale_int + 127).to(tl.uint32) << 23).to(tl.float32, bitcast=True)"- )+ _orig_src = _kern_obj._src+ _mod_src = _orig_src.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 _mod_src != _orig_src:+ _kern_obj._unsafe_update_src(_mod_src)+ _log("[setup] quantizer + kernel patched OK")+ else:+ _log("[setup] quantizer patched, kernel string mismatch")- if new_src != src:- quant_fn._unsafe_update_src(new_src)- _QUANT_PATCHED = True- print("[mm-opt] Patched _mxfp4_quant_op: integer E8M0 scale", file=sys.stderr, flush=True)- else:- # Check what patterns exist- has_log2 = "tl.log2" in src- has_floor = "tl.floor" in src- has_exp2 = "tl.exp2" in src- print(f"[mm-opt] Quant patch NO-OP: log2={has_log2}, floor={has_floor}, exp2={has_exp2}", file=sys.stderr, flush=True)- except Exception as e:- print(f"[mm-opt] Quant patch failed: {e}", file=sys.stderr, flush=True)+ _check = _quant_ref._src if hasattr(_quant_ref, '_src') else ''+ _log(f"[setup] has hw asm: {'inline_asm_elementwise' in _check}")+ except Exception as _exc:+ import traceback+ _log(f"[setup] patch FAILED: {_exc}")+ traceback.print_exc(file=_io.stderr)- def _get_jit_fn(kernel):- """Unwrap Heuristics/Autotuner wrapper to get the JITFunction with .src."""- fn = kernel- # Unwrap up to 3 levels, stopping when we find .src- for _ in range(3):- if hasattr(fn, 'src'):- return fn- if hasattr(fn, 'fn'):- fn = fn.fn- else:- break- # If no .src found, return whatever we have- return fn+ # ---- Native HIP accumulator merger for K-partitioned runs ----+ _MERGER_HIP = r"""+ #include <hip/hip_runtime.h>+ __device__ __forceinline__ unsigned short to_bf16(float val) {+ unsigned int raw;+ __builtin_memcpy(&raw, &val, sizeof(raw));+ unsigned int bias = ((raw >> 16) & 1) + 0x7FFFu;+ return (unsigned short)((raw + bias) >> 16);+ }- def _patch_gemm_kernel():- """- Patch the main GEMM kernel source to add:- - .wt store modifier (avoids L2 pollution from output writes)- - fast_math=True on tl.dot_scaled- - acc=accumulator for in-place accumulation- - eviction_policy="evict_last" on A loads- Also acts as cache-bust to force recompilation with patched quant op.- """- global _KERNEL_PATCHED- if _KERNEL_PATCHED:- return+ template <int NP>+ __global__ void vec4_sum(const float* __restrict__ src,+ unsigned short* __restrict__ dst, int len) {+ int base = (blockIdx.x * blockDim.x + threadIdx.x) * 4;+ if (base + 3 < len) {+ float4 acc = *reinterpret_cast<const float4*>(src + base);+ #pragma unroll+ for (int p = 1; p < NP; p++) {+ float4 part = *reinterpret_cast<const float4*>(src + p * len + base);+ acc.x += part.x; acc.y += part.y; acc.z += part.z; acc.w += part.w;+ }+ unsigned short a = to_bf16(acc.x), b = to_bf16(acc.y);+ unsigned short c = to_bf16(acc.z), d = to_bf16(acc.w);+ *reinterpret_cast<unsigned long long*>(dst + base) =+ (unsigned long long)a | ((unsigned long long)b << 16) |+ ((unsigned long long)c << 32) | ((unsigned long long)d << 48);+ } else {+ for (int j = base; j < len && j < base + 4; j++) {+ float acc = src[j];+ #pragma unroll+ for (int p = 1; p < NP; p++) acc += src[p * len + j];+ dst[j] = to_bf16(acc);+ }+ }+ }- if _DIRECT_KERNEL is None:- return+ __global__ void scalar_sum(const float* __restrict__ src,+ unsigned short* __restrict__ dst,+ int len, int np) {+ int gid = blockIdx.x * blockDim.x + threadIdx.x;+ if (gid < len) {+ float acc = src[gid];+ for (int p = 1; p < np; p++) acc += src[p * len + gid];+ dst[gid] = to_bf16(acc);+ }+ }- try:- jit_fn = _get_jit_fn(_DIRECT_KERNEL)- if not hasattr(jit_fn, 'src'):- print(f"[mm-opt] Kernel has no .src, type chain: {type(_DIRECT_KERNEL).__name__}", file=sys.stderr, flush=True)- # Try to find _unsafe_update_src at any level- for attr_name in ['src', '_unsafe_update_src']:- for obj in [_DIRECT_KERNEL, getattr(_DIRECT_KERNEL, 'fn', None)]:- if obj and hasattr(obj, attr_name):- print(f"[mm-opt] Found {attr_name} on {type(obj).__name__}", file=sys.stderr, flush=True)- return- src = jit_fn.src- new_src = src+ void merge_partials(torch::Tensor src, torch::Tensor dst, int R, int C, int np) {+ int len = R * C;+ const float* sp = src.data_ptr<float>();+ unsigned short* dp = reinterpret_cast<unsigned short*>(dst.data_ptr());+ const int thr = 64, stride = thr * 4;+ const int nblk = (len + stride - 1) / stride;+ switch (np) {+ case 2: vec4_sum<2><<<nblk, thr>>>(sp, dp, len); break;+ case 3: vec4_sum<3><<<nblk, thr>>>(sp, dp, len); break;+ case 4: vec4_sum<4><<<nblk, thr>>>(sp, dp, len); break;+ case 7: vec4_sum<7><<<nblk, thr>>>(sp, dp, len); break;+ case 8: vec4_sum<8><<<nblk, thr>>>(sp, dp, len); break;+ default: {+ const int t2 = 256, b2 = (len + t2 - 1) / t2;+ scalar_sum<<<b2, t2>>>(sp, dp, len, np);+ break;+ }+ }+ }+ """+ _MERGER_HDR = "void merge_partials(torch::Tensor src, torch::Tensor dst, int R, int C, int np);"- # 1. Add .wt store modifier on tl.store for y_ptr (final output)- # Match tl.store(y_ptr + ...) calls and add cache_modifier=".wt"- # Be careful not to double-add- if 'cache_modifier=".wt"' not in new_src:- new_src = re.sub(- r'(tl\.store\(\s*y_ptr\s*\+[^)]+)(,\s*mask=[^)]+)?\)',- lambda m: m.group(0).rstrip(')') + ', cache_modifier=".wt")',- new_src- )+ _HAS_HIP_MERGER = False+ try:+ from torch.utils.cpp_extension import load_inline as _jit+ _jit_t0 = _clk.time()+ _hip_merger = _jit(+ name="fp4_kmerge",+ cpp_sources=[_MERGER_HDR],+ cuda_sources=[_MERGER_HIP],+ functions=["merge_partials"],+ verbose=False,+ extra_cuda_cflags=["--offload-arch=gfx950", "-O3"],+ )+ _HAS_HIP_MERGER = True+ _log(f"[setup] HIP merger ready ({_clk.time()-_jit_t0:.1f}s)")+ except Exception as _exc:+ _log(f"[setup] HIP merger unavailable: {_exc}")- # 2. Add fast_math=True and acc=accumulator on tl.dot_scaled- if 'fast_math=True' not in new_src:- # Replace: accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")- # With: accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", acc=accumulator, fast_math=True)- new_src = re.sub(- r'accumulator\s*\+=\s*tl\.dot_scaled\(([^)]+)\)',- r'accumulator = tl.dot_scaled(\1, acc=accumulator, fast_math=True)',- new_src- )- # 3. Add eviction_policy on A loads- if 'evict_last' not in new_src:- new_src = re.sub(- r'(tl\.load\(\s*a_ptr\s*\+[^)]+)(,\s*mask=[^)]+)?\)',- lambda m: m.group(0).rstrip(')') + ', eviction_policy="evict_last")',- new_src- )+ # ---- Partition alignment utility ----+ def _align_parts(kh, bk, np):+ span = triton.cdiv((2 * triton.cdiv(kh, np)), bk) * bk+ while np > 1 and bk > 16:+ ok = (kh % (span // 2) == 0 and span % bk == 0 and kh % (bk // 2) == 0)+ if ok:+ break+ elif kh % (span // 2) != 0 and np > 1:+ np //= 2+ elif span % bk != 0:+ np = np // 2 if np > 1 else np+ if np <= 1 and bk > 16:+ bk //= 2+ elif kh % (bk // 2) != 0 and bk > 16:+ bk //= 2+ else:+ break+ span = triton.cdiv((2 * triton.cdiv(kh, np)), bk) * bk+ return span, bk, np- if new_src != src:- jit_fn._unsafe_update_src(new_src)- _KERNEL_PATCHED = True- print("[mm-opt] Patched GEMM kernel: .wt + fast_math + acc + eviction_policy", file=sys.stderr, flush=True)- except Exception as e:- print(f"[mm-opt] Kernel patch failed: {e}", file=sys.stderr, flush=True)+ # ---- Occupancy-driven config resolver ----+ _resolved = {}- def _patch_heuristics():- """Monkey-patch GRID_MN and EVEN_K heuristics to constants."""- if _DIRECT_KERNEL is None:- return- try:- if hasattr(_DIRECT_KERNEL, 'values') and 'GRID_MN' in _DIRECT_KERNEL.values:- _DIRECT_KERNEL.values['GRID_MN'] = lambda args: 1- if hasattr(_DIRECT_KERNEL, 'values') and 'EVEN_K' in _DIRECT_KERNEL.values:- _DIRECT_KERNEL.values['EVEN_K'] = lambda args: True- except Exception:- pass+ def _shape_config(batch, cols, depth):+ tag = (batch, cols, depth)+ if tag in _resolved:+ return _resolved[tag]+ kh = depth // 2-- # ── Config computation ──────────────────────────────────────────────────────-- def _get_cfg(m: int, n: int, k: int):- """- Compute per-shape config. Returns dict with all Triton kernel parameters.- Implements the tuned config from 200+ experiments.- """- cached = _SHAPE_CFG_CACHE.get((m, n, k))- if cached is not None:- return cached-- tiles_bm16_n128 = _ceil_div(m, 16) * _ceil_div(n, 128)-- # BLOCK_M selection- if m <= 32 or (m <= 128 and tiles_bm16_n128 < _LOW_UTIL_THRESHOLD):- block_m = 8+ if batch <= 32:+ bm, bn = 8, 128+ wave_est = ((batch + bm - 1) // bm) * ((cols + 127) // 128)+ np = 1+ if depth >= 4096:+ np = 7+ elif depth >= 2048:+ np = 2 if (wave_est * 2 >= (_ENGINES * 3) // 4 and wave_est * 2 <= _ENGINES) else 4+ elif depth >= 1536:+ np = 2 if (wave_est * 2 >= (_ENGINES * 3) // 4 and wave_est * 2 <= _ENGINES) else 3+ bk = 256 if depth <= np * 512 or (np == 2 and depth <= np * 1024) else 512+ if wave_est * np < (_ENGINES * 3) // 4:+ bn = 64+ total_wg = ((batch + bm - 1) // bm) * ((cols + bn - 1) // bn) * np+ wpe = 2 if total_wg > _ENGINES else 1else:- block_m = 16+ bm = 16+ if batch <= 128:+ est16 = ((batch + 15) // 16) * ((cols + 127) // 128)+ if est16 < (_ENGINES * 3) // 4:+ bm = 8+ wave_est = ((batch + bm - 1) // bm) * ((cols + 127) // 128)+ bn, np = 128, 1+ if _ENGINES // 2 <= wave_est <= _ENGINES and (depth >= 7168 or (depth >= 2048 and bm == 8)):+ np = 2+ elif wave_est < _ENGINES // 2 and depth > 512:+ if depth >= 4096:+ np = 2 if wave_est * 2 >= _ENGINES else 7+ elif depth >= 2048:+ np = 2+ elif depth >= 1536:+ np = 3+ bk = 256 if depth <= max(np * 4096, 2048) else 512+ if wave_est * np < (_ENGINES * 3) // 4:+ bn = 64+ total_wg = ((batch + bm - 1) // bm) * ((cols + bn - 1) // bn) * np+ wpe = 2 if total_wg > _ENGINES else 1- tiles_for_split = _ceil_div(m, block_m) * _ceil_div(n, 128)-- # KSPLIT routing- if m <= 32:- if k >= 4096:- ksplit = 7- elif k >= 2048:- tiles_128 = _ceil_div(m, block_m) * _ceil_div(n, 128)- if tiles_128 * 2 >= _LOW_UTIL_THRESHOLD and tiles_128 * 2 <= _CU:- ksplit = 2- else:- ksplit = 4- elif k >= 1536:- tiles_128 = _ceil_div(m, block_m) * _ceil_div(n, 128)- if tiles_128 * 2 >= _LOW_UTIL_THRESHOLD and tiles_128 * 2 <= _CU:- ksplit = 2- else:- ksplit = 3- else:- ksplit = 1- elif k >= 7168 and (_CU // 2) <= tiles_for_split <= _CU:- ksplit = 2- elif block_m == 8 and k >= 2048 and (_CU // 2) <= tiles_for_split <= _CU:- ksplit = 2- else:- ksplit = 1-- # BLOCK_K selection (BK=256 pipeline breakthrough)- if m <= 32:- if ksplit == 2 and k <= ksplit * 1024:- block_k = 256- elif k <= ksplit * 512:- block_k = 256- else:- block_k = 512- else:- if k <= max(ksplit * 4096, 2048):- block_k = 256- else:- block_k = 512-- # BLOCK_N selection- block_n = 64 if (tiles_for_split * ksplit) < _LOW_UTIL_THRESHOLD else 128-- # waves_per_eu- wgs = _ceil_div(m, block_m) * _ceil_div(n, max(block_n, 32)) * ksplit- waves_per_eu = 2 if wgs > _CU else 1-- if (m, n, k) == (16, 2112, 7168):- waves_per_eu = 2- if (m, n, k) == (64, 7168, 2048):- waves_per_eu = 1-- cfg = {- "BLOCK_SIZE_M": block_m,- "BLOCK_SIZE_N": max(block_n, 32),- "BLOCK_SIZE_K": block_k,- "GROUP_SIZE_M": 1,- "NUM_KSPLIT": ksplit,- "SPLITK_BLOCK_SIZE": max(k // max(ksplit, 1), 64),- "num_stages": 2,- "num_warps": 4,- "waves_per_eu": waves_per_eu,- "matrix_instr_nonkdim": 16,- "cache_modifier": ".cg",+ params = {+ "BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": max(bn, 32), "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": np,}- _SHAPE_CFG_CACHE[(m, n, k)] = cfg- return cfg+ if params["NUM_KSPLIT"] > 1:+ span, bk2, np2 = _align_parts(kh, params["BLOCK_SIZE_K"], params["NUM_KSPLIT"])+ params["SPLITK_BLOCK_SIZE"] = span+ params["BLOCK_SIZE_K"] = bk2+ params["NUM_KSPLIT"] = np2+ if params["BLOCK_SIZE_K"] >= 2 * kh:+ params["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * kh)+ params["SPLITK_BLOCK_SIZE"] = 2 * kh+ params["NUM_KSPLIT"] = 1+ params["BLOCK_SIZE_N"] = max(params["BLOCK_SIZE_N"], 32)- def _shape_uses_disable_lsr(m: int, k: int) -> bool:- return not (m <= 32 and k >= 1536)+ if params["NUM_KSPLIT"] == 1:+ params["SPLITK_BLOCK_SIZE"] = 2 * kh+ real_np, padded_np = None, None+ if params["NUM_KSPLIT"] > 1:+ real_np = triton.cdiv(kh, params["SPLITK_BLOCK_SIZE"] // 2)+ padded_np = triton.next_power_of_2(params["NUM_KSPLIT"])- def _set_disable_lsr(enabled: bool):- previous = os.environ.get("DISABLE_LLVM_OPT")- if enabled:- os.environ["DISABLE_LLVM_OPT"] = "disable-lsr"- else:- os.environ.pop("DISABLE_LLVM_OPT", None)- return previous+ m_tiles = triton.cdiv(batch, params["BLOCK_SIZE_M"])+ n_tiles = triton.cdiv(cols, params["BLOCK_SIZE_N"])+ launch_grid = (params["NUM_KSPLIT"] * m_tiles * n_tiles,)+ red_grid = None+ if params["NUM_KSPLIT"] > 1:+ red_grid = (triton.cdiv(batch, 16), triton.cdiv(cols, 16))+ bundle = (+ params, real_np, padded_np, launch_grid, red_grid,+ kh, params["BLOCK_SIZE_M"], params["BLOCK_SIZE_N"],+ params["BLOCK_SIZE_K"], params["NUM_KSPLIT"],+ params["SPLITK_BLOCK_SIZE"], params["waves_per_eu"],+ )+ _resolved[tag] = bundle+ return bundle- def _restore_disable_lsr(previous):- if previous is None:- os.environ.pop("DISABLE_LLVM_OPT", None)- else:- os.environ["DISABLE_LLVM_OPT"] = previous+ # ---- Phased JIT warmup ----+ _warm_t0 = _clk.time()+ _warmed = {}+ _phase1 = {}+ _phase2 = {}+ _red_set = set()- # ── Pre-shuffled B views ────────────────────────────────────────────────────+ for _dn, _dk in _DIM_PAIRS:+ for _db in _BATCH_DIMS:+ _cb, _rn, _rp, _, _, _, _, _, _, _, _, _ = _shape_config(_db, _dn, _dk)+ _sig = (+ _cb["BLOCK_SIZE_M"], _cb["BLOCK_SIZE_N"], _cb["BLOCK_SIZE_K"],+ _cb["NUM_KSPLIT"], _cb["SPLITK_BLOCK_SIZE"], _cb["waves_per_eu"],+ )+ if _db <= 32 and _dk >= 1536:+ _phase1.setdefault(_sig, True)+ else:+ _phase2.setdefault(_sig, True)+ if _rn is not None:+ _red_set.add((_rn, _rp))- def _get_preshuffle_views(b_shuffle, b_scale_sh, n, k):- key = (b_shuffle.data_ptr(), b_scale_sh.data_ptr(), n, k)- cached = _PRESHUFFLE_CACHE.get(key)- if cached is not None:- b_ref, s_ref, b_ps_u8, s_ps_u8 = cached- if b_ref() is b_shuffle and s_ref() is b_scale_sh:- return b_ps_u8, s_ps_u8+ for _s in _phase1:+ _phase2.pop(_s, None)- b_ps_u8 = _view_dtype(b_shuffle, torch.uint8).contiguous().view(n // 16, k * 8).contiguous()- scale_u8 = _view_dtype(b_scale_sh, torch.uint8).contiguous()- s_ps_u8 = scale_u8[:n, :(k // 32)].contiguous().view(n // 32, k).contiguous()+ _log(f"[warm] {len(_phase1)} phase1 + {len(_phase2)} phase2, {len(_red_set)} reducers")- _PRESHUFFLE_CACHE[key] = (weakref.ref(b_shuffle), weakref.ref(b_scale_sh), b_ps_u8, s_ps_u8)- return b_ps_u8, s_ps_u8+ _dummy_x = torch.zeros(32, 8192, dtype=torch.bfloat16, device="cuda")+ _dummy_w = torch.zeros(16, 65536, dtype=torch.uint8, device="cuda")+ _dummy_s = torch.zeros(16, 65536, dtype=torch.uint8, device="cuda")+ _dummy_pp = torch.zeros(16, 32, 256, dtype=torch.float32, device="cuda")+ _dummy_y = torch.zeros(32, 256, dtype=torch.bfloat16, device="cuda")+ def _fire_warmup(bm, bn, bk, ks, spk, wpe):+ cfg = {+ "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,+ }+ target = _dummy_pp if ks > 1 else _dummy_y+ _gemm_a16wfp4_preshuffle_kernel[(max(ks, 1),)](+ _dummy_x, _dummy_w, target, _dummy_s, bm, bn, spk // 2,+ _dummy_x.stride(0), _dummy_x.stride(1),+ _dummy_w.stride(0), _dummy_w.stride(1),+ 0 if ks <= 1 else _dummy_pp.stride(0),+ _dummy_y.stride(0) if ks <= 1 else _dummy_pp.stride(1),+ _dummy_y.stride(1) if ks <= 1 else _dummy_pp.stride(2),+ _dummy_s.stride(0), _dummy_s.stride(1),+ PREQUANT=True, **cfg,+ )- def _get_output(m, n):- key = (m, n)- out = _OUT_CACHE.get(key)- if out is None or out.shape != (m, n):- out = torch.empty((m, n), dtype=torch.bfloat16, device="cuda")- _OUT_CACHE[key] = out- return out+ _log("[warm] phase 1 (no lsr)...")+ for _sig in sorted(_phase1):+ try:+ _fire_warmup(*_sig)+ _warmed[_sig] = 1+ _log(f" {_sig[0]}x{_sig[1]}x{_sig[2]} ks={_sig[3]} ({_clk.time()-_warm_t0:.0f}s)")+ except Exception as _exc:+ _log(f" {_sig}: ERR {_exc}")+ _env.environ["DISABLE_LLVM_OPT"] = "disable-lsr"+ _log(f"[warm] phase 2 (with lsr) @ {_clk.time()-_warm_t0:.0f}s")- def _get_partials(num_ksplit, m, n):- key = (num_ksplit, m, n)- p = _PARTIAL_CACHE.get(key)- if p is None or p.shape != (num_ksplit, m, n):- p = torch.empty((num_ksplit, m, n), dtype=torch.float32, device="cuda")- _PARTIAL_CACHE[key] = p- return p--- # ── Runtime resolution ──────────────────────────────────────────────────────-- def _resolve_runtime():- global _INIT_DONE, _DIRECT_KERNEL, _REDUCE_KERNEL, _GET_SPLITK- if _INIT_DONE:- return- _INIT_DONE = True-+ for _idx, _sig in enumerate(sorted(_phase2)):+ if _clk.time() - _warm_t0 > 200:+ _log(f" timeout, {len(_phase2) - _idx} skipped")+ breaktry:- kernel_mod = importlib.import_module("aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4")- _DIRECT_KERNEL = getattr(kernel_mod, "_gemm_a16wfp4_preshuffle_kernel", None)- except Exception:- pass+ _fire_warmup(*_sig)+ _warmed[_sig] = 2+ _log(f" {_sig[0]}x{_sig[1]}x{_sig[2]} ks={_sig[3]} ({_clk.time()-_warm_t0:.0f}s)")+ except Exception as _exc:+ _log(f" {_sig}: ERR {_exc}")+ _log(f"[warm] reducers...")+ for _rn, _rp in sorted(_red_set):+ if _clk.time() - _warm_t0 > 230:+ _log(" timeout")+ breaktry:- reduce_mod = importlib.import_module("aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4")- _REDUCE_KERNEL = getattr(reduce_mod, "_gemm_afp4wfp4_reduce_kernel", None)+ _gemm_afp4wfp4_reduce_kernel[(1, 1)](+ _dummy_pp, _dummy_y, 16, 16,+ _dummy_pp.stride(0), _dummy_pp.stride(1), _dummy_pp.stride(2),+ _dummy_y.stride(0), _dummy_y.stride(1), 16, 16, _rn, _rp,+ )except Exception:pass- try:- splitk_mod = importlib.import_module("aiter.ops.triton.gemm.basic.gemm_afp4wfp4")- _GET_SPLITK = getattr(splitk_mod, "get_splitk", None)- except Exception:- pass+ del _dummy_x, _dummy_w, _dummy_s, _dummy_pp, _dummy_y, _fire_warmup+ del _phase1, _phase2, _red_set+ torch.cuda.empty_cache()+ _log(f"[warm] done: {len(_warmed)} configs in {_clk.time()-_warm_t0:.0f}s")- # Apply patches- _patch_heuristics()- _patch_quant_op()- _patch_gemm_kernel()+ _mem.disable()- # Nuclear pre-warming- _prewarm_all()+ # ---- Runtime dispatch state ----+ _wt_cache = {}+ _dest_buf = {}+ _frag_buf = {}+ _seen = set()- def _finalize_cfg(cfg, k):- """Apply _get_splitk alignment and fix up config for kernel call."""- cfg = dict(cfg)- if cfg["NUM_KSPLIT"] > 1 and _GET_SPLITK is not None:- 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"] = int(triton.next_power_of_2(2 * k))- cfg["SPLITK_BLOCK_SIZE"] = 2 * k- cfg["NUM_KSPLIT"] = 1+ def _prepare_wt(data):+ addr = data[3].data_ptr()+ if addr not in _wt_cache:+ n_dim = data[3].shape[0]+ k_bytes = data[3].shape[1]+ sr, sc = data[4].shape+ n_grp = n_dim // 32+ w_view = data[3].view(torch.uint8).reshape(n_dim // 16, k_bytes * 16)+ s_view = data[4].view(torch.uint8).reshape(sr // 32, sc * 32)[:n_grp].contiguous()+ _wt_cache[addr] = (w_view, s_view, w_view.stride(0), s_view.stride(0))+ return _wt_cache[addr]- cfg["BLOCK_SIZE_N"] = max(cfg["BLOCK_SIZE_N"], 32)- if cfg["NUM_KSPLIT"] <= 1:- cfg["NUM_KSPLIT"] = 1- cfg["SPLITK_BLOCK_SIZE"] = 2 * k- return cfg-- # ── Pre-warming ─────────────────────────────────────────────────────────────-- def _prewarm_all():- """Nuclear pre-warming with selective disable-lsr."""- if _DIRECT_KERNEL is None:- return-- all_m = [1, 2, 4, 8, 16, 32, 64, 128, 256]- all_n = [2112, 2880, 3072, 4096, 7168]- all_k = [512, 1536, 2048, 7168]-- phase1_cfgs = set() # no disable-lsr (M<=32 K>=1536)- phase2_cfgs = set() # with disable-lsr (everything else)-- for m in all_m:- for n in all_n:- for k in all_k:- cfg = _get_cfg(m, n, k)- final = _finalize_cfg(cfg, k)- key = (- final["BLOCK_SIZE_M"], final["BLOCK_SIZE_N"],- final["BLOCK_SIZE_K"], final["NUM_KSPLIT"],- final["SPLITK_BLOCK_SIZE"], final["num_stages"],- final["num_warps"], final["waves_per_eu"],- )- if _shape_uses_disable_lsr(m, k):- phase2_cfgs.add(key)- else:- phase1_cfgs.add(key)-- # Phase 1: compile without disable-lsr- prev = _set_disable_lsr(False)- _prewarm_configs(phase1_cfgs)- _restore_disable_lsr(prev)-- # Phase 2: compile with disable-lsr- prev = _set_disable_lsr(True)- _prewarm_configs(phase2_cfgs)- _restore_disable_lsr(prev)-- # Pre-warm reduce kernel- if _REDUCE_KERNEL is not None:- _prewarm_reduce()-- print(f"[mm-opt] Pre-warmed {len(phase1_cfgs)} no-lsr + {len(phase2_cfgs)} lsr configs",- file=sys.stderr, flush=True)--- def _prewarm_configs(cfg_keys):- if _DIRECT_KERNEL is None:- return- for bm, bn, bk, ks, spk, stages, warps, wpe in cfg_keys:- try:- test_m, test_n, test_k = bm, bn, max(bk, 256)- a = torch.zeros((test_m, test_k), dtype=torch.bfloat16, device="cuda")- b_w = torch.zeros((test_n // 16, test_k * 8), dtype=torch.uint8, device="cuda")- b_s = torch.zeros((test_n // 32, test_k), dtype=torch.uint8, device="cuda")- if ks > 1:- out = torch.zeros((ks, test_m, test_n), dtype=torch.float32, device="cuda")- else:- out = torch.zeros((test_m, test_n), dtype=torch.bfloat16, device="cuda")-- grid = lambda meta: (ks * _ceil_div(test_m, bm) * _ceil_div(test_n, bn),)- _DIRECT_KERNEL[grid](- a, b_w, out, b_s,- test_m, test_n, test_k,- test_k, 1, test_k * 8, 1,- 0 if ks <= 1 else test_m * test_n,- test_n, 1, test_k, 1,- PREQUANT=True,- BLOCK_SIZE_M=bm, BLOCK_SIZE_N=bn, BLOCK_SIZE_K=bk,- GROUP_SIZE_M=1, NUM_KSPLIT=ks, SPLITK_BLOCK_SIZE=spk,- num_stages=stages, num_warps=warps, waves_per_eu=wpe,- matrix_instr_nonkdim=16, cache_modifier=".cg",- )- except Exception:- pass--- def _prewarm_reduce():- if _REDUCE_KERNEL is None:- return- for ks in [2, 3, 4, 7, 8]:- try:- y_pp = torch.zeros((ks, 16, 128), dtype=torch.float32, device="cuda")- y = torch.zeros((16, 128), dtype=torch.bfloat16, device="cuda")- nk_pow2 = int(triton.next_power_of_2(ks))- grid_r = (_ceil_div(16, 16), _ceil_div(128, 16))- _REDUCE_KERNEL[grid_r](- y_pp, y, 16, 128,- 16 * 128, 128, 1, 128, 1,- 16, 16, ks, nk_pow2,- )- except Exception:- pass--- # ── Fallback ────────────────────────────────────────────────────────────────-- def _quant_ref(x):- x_fp4, raw_scale = dynamic_mxfp4_quant(x)- scale_sh = e8m0_shuffle(raw_scale)- return x_fp4.view(dtypes.fp4x2), scale_sh.view(dtypes.fp8_e8m0)--- def _run_fallback_gemm(a, b_shuffle, a_scale_sh, b_scale_sh):- return aiter.gemm_a4w4(a, b_shuffle, a_scale_sh, b_scale_sh, dtype=dtypes.bf16, bpreshuffle=True)--- # ── Main dispatch ───────────────────────────────────────────────────────────-- @torch.inference_mode()def custom_kernel(data: input_t) -> output_t:- A, _B, _B_q, B_shuffle, B_scale_sh = data- if not A.is_contiguous():- A = A.contiguous()+ X = data[0]+ if not X.is_contiguous():+ X = X.contiguous()+ ndims = X.ndim+ X_flat = X if ndims == 2 else X.view(-1, X.shape[-1])+ batch = X_flat.shape[0]+ cols = data[3].shape[0]+ depth = data[3].shape[1] * 2- m, k = A.shape- n = B_shuffle.shape[0]- shape = (m, n, k)+ (params, real_np, padded_np, launch_grid, red_grid,+ kh, bm, bn, bk, ks, spk, wpe) = _shape_config(batch, cols, depth)- _resolve_runtime()+ tag = (batch, cols, depth)+ if tag not in _seen:+ _seen.add(tag)+ _log(f"[run] {batch}x{cols}x{depth} bm={bm} bn={bn} bk={bk} ks={ks} wpe={wpe}")- if _DIRECT_KERNEL is None:- A_q, A_scale_sh = _quant_ref(A)- return _run_fallback_gemm(A_q, B_shuffle, A_scale_sh, B_scale_sh)+ okey = (batch, cols)+ if okey not in _dest_buf:+ _dest_buf[okey] = torch.empty((batch, cols), dtype=torch.bfloat16, device="cuda")+ dest = _dest_buf[okey]- b_ps_u8, s_ps_u8 = _get_preshuffle_views(B_shuffle, B_scale_sh, n, k)- runtime_n = b_ps_u8.shape[0] * 16- runtime_k = b_ps_u8.shape[1] // 16+ w_view, s_view, sw0, ss0 = _prepare_wt(data)- cfg = _get_cfg(m, n, k)-- # Set disable-lsr based on shape- if _shape_uses_disable_lsr(m, k):- os.environ["DISABLE_LLVM_OPT"] = "disable-lsr"+ if ks > 1:+ fkey = (padded_np, batch, cols)+ if fkey not in _frag_buf:+ _frag_buf[fkey] = torch.empty(+ (padded_np, batch, cols), dtype=torch.float32, device="cuda"+ )+ frags = _frag_buf[fkey]+ stride_p, stride_r = batch * cols, colselse:- os.environ.pop("DISABLE_LLVM_OPT", None)+ frags = None+ stride_p, stride_r = 0, cols- final = _finalize_cfg(cfg, runtime_k)-- num_ksplit = final["NUM_KSPLIT"]- bm = final["BLOCK_SIZE_M"]- bn = final["BLOCK_SIZE_N"]-- y = _get_output(m, runtime_n)-- if num_ksplit > 1:- y_pp = _get_partials(num_ksplit, m, runtime_n)- out = y_pp- else:- y_pp = None- out = y-- # Pre-computed strides (all contiguous)- # A is (m, k) contiguous → stride(0) = k (original K, not runtime_k)- stride_a0, stride_a1 = k, 1- stride_bw0, stride_bw1 = b_ps_u8.shape[1], 1- stride_bs0, stride_bs1 = s_ps_u8.shape[1], 1-- if y_pp is not None:- stride_ypp0 = m * runtime_n- stride_y0, stride_y1 = runtime_n, 1- else:- stride_ypp0 = 0- stride_y0, stride_y1 = runtime_n, 1-- grid = lambda meta: (- meta["NUM_KSPLIT"] * _ceil_div(m, int(meta["BLOCK_SIZE_M"])) * _ceil_div(runtime_n, int(meta["BLOCK_SIZE_N"])),+ _gemm_a16wfp4_preshuffle_kernel[launch_grid](+ X_flat, w_view,+ dest if frags is None else frags,+ s_view, batch, cols, kh,+ depth, 1, sw0, 1,+ stride_p, stride_r, 1,+ ss0, 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,)- try:- _DIRECT_KERNEL[grid](- A, b_ps_u8, out, s_ps_u8,- m, runtime_n, runtime_k,- stride_a0, stride_a1,- stride_bw0, stride_bw1,- stride_ypp0,- stride_y0, stride_y1,- stride_bs0, stride_bs1,- PREQUANT=True,- **final,- )-- if y_pp is not None:- actual_ksplit = int(triton.cdiv(runtime_k, int(final["SPLITK_BLOCK_SIZE"]) // 2))- # Triton reduce- nk_pow2 = int(triton.next_power_of_2(int(final["NUM_KSPLIT"])))- grid_r = (_ceil_div(m, 16), _ceil_div(runtime_n, 16))- _REDUCE_KERNEL[grid_r](- y_pp, y, m, runtime_n,- y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),- y.stride(0), y.stride(1),- 16, 16, actual_ksplit, nk_pow2,+ if frags is not None:+ if _HAS_HIP_MERGER:+ _hip_merger.merge_partials(frags, dest, batch, cols, real_np)+ else:+ _gemm_afp4wfp4_reduce_kernel[red_grid](+ frags, dest, batch, cols,+ batch * cols, cols, 1, cols, 1,+ 16, 16, real_np, padded_np,)- if shape not in _LOGGED_PATHS:- _LOGGED_PATHS.add(shape)- bk = final["BLOCK_SIZE_K"]- ks = final["NUM_KSPLIT"]- wpe = final["waves_per_eu"]- print(f"[mm-opt] shape={shape} bm={bm},bn={bn},bk={bk},ks={ks},wpe={wpe}",- file=sys.stderr, flush=True)-- return y-- except Exception as e:- if shape not in _LOGGED_PATHS:- _LOGGED_PATHS.add(shape)- print(f"[mm-opt] shape={shape} FALLBACK: {e}", file=sys.stderr, flush=True)- A_q, A_scale_sh = _quant_ref(A)- return _run_fallback_gemm(A_q, B_shuffle, A_scale_sh, B_scale_sh)+ return dest if ndims == 2 else dest.view(*X.shape[:-1], cols)
scrolls · 1083 diff lines total
Best evidence level for this revision: reported
JSON