submission 738857
.jonnss · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 631 lines, June 9 Researcher Reciprocity License v1.0.
Submission_v01.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-738857?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:be8340e203b1408c32cf980a1a2dfa9964f3aaa3735013d8252bd09d28f5681e
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
Optimized MXFP4 GEMM submission based on GEMM-Reference.md best practices.split-k
_GET_SPLITK = Nonetile-k = 256
5. Tuned per-shape configs (BK=256 pipeline, KSPLIT routing, BM=8 for M<=32)tile-m = 8
5. Tuned per-shape configs (BK=256 pipeline, KSPLIT routing, BM=8 for M<=32)Kernel source
Submission_v01.py631 lines
"""
Optimized MXFP4 GEMM submission based on GEMM-Reference.md best practices.
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
# ── 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 aiter
import torch
import triton
import 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 task import input_t, output_t
# ── Global state ────────────────────────────────────────────────────────────
_CU = 256
_LOW_UTIL_THRESHOLD = (_CU * 3) // 4
_PRESHUFFLE_CACHE = {}
_OUT_CACHE = {}
_PARTIAL_CACHE = {}
_SHAPE_CFG_CACHE = {}
_LOGGED_PATHS = set()
_DIRECT_KERNEL = None
_REDUCE_KERNEL = None
_GET_SPLITK = None
_INIT_DONE = False
_QUANT_PATCHED = False
_KERNEL_PATCHED = False
gc.disable()
torch.set_grad_enabled(False)
sys.setswitchinterval(1.0)
# ── Helpers ─────────────────────────────────────────────────────────────────
def _ceil_div(a: int, b: int) -> int:
return (a + b - 1) // b
def _view_dtype(tensor: torch.Tensor, dtype) -> torch.Tensor:
if tensor.dtype == dtype:
return tensor
return tensor.view(dtype)
# ── Source patching ─────────────────────────────────────────────────────────
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
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
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
src = quant_fn.src
pass # src loaded
new_src = src
# 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)"
)
# 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)"
)
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)
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
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
if _DIRECT_KERNEL is None:
return
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
# 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
)
# 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
)
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)
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
# ── 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
else:
block_m = 16
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",
}
_SHAPE_CFG_CACHE[(m, n, k)] = cfg
return cfg
def _shape_uses_disable_lsr(m: int, k: int) -> bool:
return not (m <= 32 and k >= 1536)
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
def _restore_disable_lsr(previous):
if previous is None:
os.environ.pop("DISABLE_LLVM_OPT", None)
else:
os.environ["DISABLE_LLVM_OPT"] = previous
# ── Pre-shuffled B views ────────────────────────────────────────────────────
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
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()
_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
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
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
try:
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
try:
reduce_mod = importlib.import_module("aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4")
_REDUCE_KERNEL = getattr(reduce_mod, "_gemm_afp4wfp4_reduce_kernel", None)
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
# Apply patches
_patch_heuristics()
_patch_quant_op()
_patch_gemm_kernel()
# Nuclear pre-warming
_prewarm_all()
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
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()
m, k = A.shape
n = B_shuffle.shape[0]
shape = (m, n, k)
_resolve_runtime()
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)
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
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"
else:
os.environ.pop("DISABLE_LLVM_OPT", None)
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"])),
)
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 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)
scrolls · 631 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 725090.
"""- Non-HIP exp-119-lite reconstruction.+ Optimized MXFP4 GEMM submission based on GEMM-Reference.md best practices.- This keeps the `v409` direct-path runtime bundle and adds stride caching from- the old exp-116 line while staying on the legal Triton reduce surface:- - flatten the direct Triton `EVEN_K` / `GRID_MN` heuristics to constants- - disable GC and autograd globally- - reduce Python thread-switch churn- - precompute hot-path contiguous strides- - replace `**cfg` launch unpacking with explicit kwargs+ 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 gcimport importlibimport os+ import reimport sysimport weakref+ # ── 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 aiterimport torch+ import triton+ import triton.language as tlfrom aiter import dtypesfrom aiter.ops.triton.quant import dynamic_mxfp4_quantfrom aiter.utility.fp4_utils import e8m0_shufflefrom task import input_t, output_t-+ # ── Global state ────────────────────────────────────────────────────────────_CU = 256_LOW_UTIL_THRESHOLD = (_CU * 3) // 4- _MAX_CACHE_ENTRIES = 16- _CUDA_DEVICE = "cuda"- _A_QUANT_CACHE = {}_PRESHUFFLE_CACHE = {}- _SHAPE_CACHE = {}_OUT_CACHE = {}_PARTIAL_CACHE = {}+ _SHAPE_CFG_CACHE = {}+ _LOGGED_PATHS = set()- _DIRECT_INIT_DONE = False- _DIRECT_HELPER = None- _DIRECT_HELPER_ACCEPTS_DICT = True- _SERIALIZE_DICT = None_DIRECT_KERNEL = None_REDUCE_KERNEL = None_GET_SPLITK = None- _TRITON = None- _DIRECT_KERNEL_SHAPE_SUPPORT = {}- _DIRECT_HELPER_SHAPE_SUPPORT = {}- _LOGGED_PATHS = {}- _DIRECT_HEURISTICS_PATCHED = False+ _INIT_DONE = False+ _QUANT_PATCHED = False+ _KERNEL_PATCHED = False- os.environ["DISABLE_LLVM_OPT"] = "disable-lsr"-gc.disable()torch.set_grad_enabled(False)sys.setswitchinterval(1.0)+ # ── Helpers ─────────────────────────────────────────────────────────────────def _ceil_div(a: int, b: int) -> int:return (a + b - 1) // b⋯ 4 unchanged linesreturn tensor.view(dtype)- def _trim_cache(cache: dict) -> None:- while len(cache) > _MAX_CACHE_ENTRIES:- cache.pop(next(iter(cache)))+ # ── Source patching ─────────────────────────────────────────────────────────+ 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- def _shape_uses_disable_lsr(m: int, k: int) -> bool:- return not (m <= 32 and k >= 1536)+ 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+ 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- def _set_disable_lsr(enabled: bool) -> str | None:- 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+ src = quant_fn.src+ pass # src loaded+ new_src = src- def _restore_disable_lsr(previous: str | None) -> None:- if previous is None:- os.environ.pop("DISABLE_LLVM_OPT", None)- else:- os.environ["DISABLE_LLVM_OPT"] = previous+ # 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)"+ )+ # 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)"+ )- def _quant_ref(x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:- 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)+ 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)- def _get_cached_a_quant(a: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:- key = a.data_ptr()- cached = _A_QUANT_CACHE.get(key)- if cached is not None:- a_ref, a_ptr, a_version, a_q, a_scale_sh = cached- if a_ref() is a and a_ptr == a.data_ptr() and a_version == a._version:- return a_q, a_scale_sh+ 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- a_q, a_scale_sh = _quant_ref(a)- _A_QUANT_CACHE[key] = (weakref.ref(a), a.data_ptr(), a._version, a_q, a_scale_sh)- stale_keys = [cache_key for cache_key, entry in _A_QUANT_CACHE.items() if entry[0]() is None]- for stale_key in stale_keys:- _A_QUANT_CACHE.pop(stale_key, None)- _trim_cache(_A_QUANT_CACHE)- return a_q, a_scale_sh+ 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- def _get_cached_preshuffle_views(- b_shuffle: torch.Tensor,- b_scale_sh: torch.Tensor,- n: int,- k: int,- ) -> tuple[torch.Tensor, torch.Tensor]:- 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_ptr, s_ptr, b_version, s_version, b_ps_u8, s_ps_u8 = cached- if (- b_ref() is b_shuffle- and s_ref() is b_scale_sh- and b_ptr == b_shuffle.data_ptr()- and s_ptr == b_scale_sh.data_ptr()- and b_version == b_shuffle._version- and s_version == b_scale_sh._version- ):- return b_ps_u8, s_ps_u8+ if _DIRECT_KERNEL is None:+ return- 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()+ 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- _PRESHUFFLE_CACHE[key] = (- weakref.ref(b_shuffle),- weakref.ref(b_scale_sh),- b_shuffle.data_ptr(),- b_scale_sh.data_ptr(),- b_shuffle._version,- b_scale_sh._version,- b_ps_u8,- s_ps_u8,- )- stale_keys = [- cache_key- for cache_key, entry in _PRESHUFFLE_CACHE.items()- if entry[0]() is None or entry[1]() is None- ]- for stale_key in stale_keys:- _PRESHUFFLE_CACHE.pop(stale_key, None)- _trim_cache(_PRESHUFFLE_CACHE)- return b_ps_u8, s_ps_u8+ # 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+ )+ # 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+ )- def _get_cached_output(device: torch.device, m: int, n: int) -> torch.Tensor:- key = (m, n)- out = _OUT_CACHE.get(key)- if out is None or out.device != device or out.shape != (m, n):- out = torch.empty((m, n), dtype=torch.bfloat16, device=_CUDA_DEVICE)- _OUT_CACHE[key] = out- _trim_cache(_OUT_CACHE)- return out+ # 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+ )+ 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)- def _get_cached_partials(device: torch.device, num_ksplit: int, m: int, n: int) -> torch.Tensor:- key = (num_ksplit, m, n)- partials = _PARTIAL_CACHE.get(key)- if partials is None or partials.device != device or partials.shape != (num_ksplit, m, n):- partials = torch.empty((num_ksplit, m, n), dtype=torch.float32, device=_CUDA_DEVICE)- _PARTIAL_CACHE[key] = partials- _trim_cache(_PARTIAL_CACHE)- return partials-- def _config_to_dict(base) -> dict[str, object]:- if base is None:- return {}- if isinstance(base, dict):- return dict(base)- kwargs = getattr(base, "kwargs", None)- if kwargs is not None:- cfg = dict(kwargs)- for attr in ("num_warps", "num_stages", "num_ctas", "waves_per_eu", "maxnreg"):- val = getattr(base, attr, None)- if val is not None:- cfg[attr] = val- return cfg- try:- return dict(base)- except Exception:- return {}--- def _resolve_runtime() -> None:- global _DIRECT_INIT_DONE, _DIRECT_HELPER, _DIRECT_HELPER_ACCEPTS_DICT, _SERIALIZE_DICT- global _DIRECT_KERNEL, _REDUCE_KERNEL, _GET_SPLITK, _TRITON, _DIRECT_HEURISTICS_PATCHED- if _DIRECT_INIT_DONE:+ def _patch_heuristics():+ """Monkey-patch GRID_MN and EVEN_K heuristics to constants."""+ if _DIRECT_KERNEL is None:return- _DIRECT_INIT_DONE = True-try:- utils_mod = importlib.import_module("aiter.ops.triton.utils.common_utils")- _SERIALIZE_DICT = getattr(utils_mod, "serialize_dict", None)+ 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: Trueexcept Exception:- _SERIALIZE_DICT = None+ pass- try:- _TRITON = importlib.import_module("triton")- except Exception:- _TRITON = None- try:- 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:- _DIRECT_KERNEL = None+ # ── Config computation ──────────────────────────────────────────────────────- try:- reduce_mod = importlib.import_module("aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4")- _REDUCE_KERNEL = getattr(reduce_mod, "_gemm_afp4wfp4_reduce_kernel", None)- except Exception:- _REDUCE_KERNEL = None-- try:- splitk_mod = importlib.import_module("aiter.ops.triton.gemm.basic.gemm_afp4wfp4")- _GET_SPLITK = getattr(splitk_mod, "get_splitk", None)- except Exception:- _GET_SPLITK = None-- candidates = []- for module_name in (- "aiter.ops.triton.gemm.basic.gemm_a16wfp4",- "aiter.ops.triton.gemm.gemm_a16wfp4",- "aiter.ops.triton.gemm.basic",- "aiter.ops.triton.gemm",- ):- try:- mod = importlib.import_module(module_name)- except Exception:- continue- candidates.extend(- [- (mod, "gemm_a16wfp4_preshuffle_", False),- (mod, "gemm_a16wfp4_preshuffle", True),- ]- )- candidates.extend(- [- (aiter, "gemm_a16wfp4_preshuffle_", False),- (aiter, "gemm_a16wfp4_preshuffle", True),- ]- )-- for holder, name, accepts_dict in candidates:- fn = getattr(holder, name, None)- if callable(fn):- _DIRECT_HELPER = fn- _DIRECT_HELPER_ACCEPTS_DICT = accepts_dict- break-- if not _DIRECT_HEURISTICS_PATCHED and _DIRECT_KERNEL is not None:- values = getattr(_DIRECT_KERNEL, "values", None)- if isinstance(values, dict):- if "EVEN_K" in values:- values["EVEN_K"] = lambda args: True- if "GRID_MN" in values:- values["GRID_MN"] = lambda args: 1- _DIRECT_HEURISTICS_PATCHED = True--- def _pick_shape_entry(m: int, n: int, k: int) -> dict[str, object]:- shape = (m, n, k)- cached = _SHAPE_CACHE.get(shape)+ 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 cachedtiles_bm16_n128 = _ceil_div(m, 16) * _ceil_div(n, 128)++ # BLOCK_M selectionif m <= 32 or (m <= 128 and tiles_bm16_n128 < _LOW_UTIL_THRESHOLD):block_m = 8else:block_m = 16tiles_for_split = _ceil_div(m, block_m) * _ceil_div(n, 128)++ # KSPLIT routingif m <= 32:if k >= 4096:ksplit = 7elif k >= 2048:- ksplit = 4+ 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 = 4elif k >= 1536:- ksplit = 3+ 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 = 3else:ksplit = 1- elif k >= 2048 and tiles_for_split > _CU and tiles_for_split <= (_CU * 3) // 2:- ksplit = 2elif k >= 7168 and (_CU // 2) <= tiles_for_split <= _CU:ksplit = 2elif block_m == 8 and k >= 2048 and (_CU // 2) <= tiles_for_split <= _CU:⋯ 1 unchanged lineselse:ksplit = 1- block_k = 256 if k <= (ksplit * 512) else 512+ # 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 selectionblock_n = 64 if (tiles_for_split * ksplit) < _LOW_UTIL_THRESHOLD else 128- wgs = _ceil_div(m, block_m) * _ceil_div(n, block_n) * ksplit++ # 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": block_n,+ "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": 2 if wgs > _CU else 1,+ "waves_per_eu": waves_per_eu,"matrix_instr_nonkdim": 16,"cache_modifier": ".cg",}- if shape == (16, 2112, 7168):- cfg["waves_per_eu"] = 2- if shape == (64, 7168, 2048):- cfg["waves_per_eu"] = 1- entry = {"cfg": cfg}- _SHAPE_CACHE[shape] = entry- _trim_cache(_SHAPE_CACHE)- return entry+ _SHAPE_CFG_CACHE[(m, n, k)] = cfg+ return cfg- def _prepare_helper_cfg(m: int, n: int, k: int) -> dict[str, object]:- cfg = dict(_config_to_dict(_pick_shape_entry(m, n, k)["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 _TRITON is not None and 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 _shape_uses_disable_lsr(m: int, k: int) -> bool:+ return not (m <= 32 and k >= 1536)- 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+ 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- def _prepare_direct_cfg(m: int, n: int, k: int, runtime_k: int) -> dict[str, object]:- cfg = dict(_config_to_dict(_pick_shape_entry(m, n, k)["cfg"]))++ def _restore_disable_lsr(previous):+ if previous is None:+ os.environ.pop("DISABLE_LLVM_OPT", None)+ else:+ os.environ["DISABLE_LLVM_OPT"] = previous+++ # ── Pre-shuffled B views ────────────────────────────────────────────────────++ 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++ 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()++ _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+++ 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+++ 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++ try:+ 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++ try:+ reduce_mod = importlib.import_module("aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4")+ _REDUCE_KERNEL = getattr(reduce_mod, "_gemm_afp4wfp4_reduce_kernel", None)+ 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++ # Apply patches+ _patch_heuristics()+ _patch_quant_op()+ _patch_gemm_kernel()++ # Nuclear pre-warming+ _prewarm_all()+++ 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(- runtime_k, cfg["BLOCK_SIZE_K"], cfg["NUM_KSPLIT"]+ k, cfg["BLOCK_SIZE_K"], cfg["NUM_KSPLIT"])cfg["SPLITK_BLOCK_SIZE"] = splitk_block_sizecfg["BLOCK_SIZE_K"] = block_size_kcfg["NUM_KSPLIT"] = num_ksplit- if _TRITON is not None and cfg["BLOCK_SIZE_K"] >= 2 * runtime_k:- cfg["BLOCK_SIZE_K"] = int(_TRITON.next_power_of_2(2 * runtime_k))- cfg["SPLITK_BLOCK_SIZE"] = 2 * runtime_k+ if cfg["BLOCK_SIZE_K"] >= 2 * k:+ cfg["BLOCK_SIZE_K"] = int(triton.next_power_of_2(2 * k))+ cfg["SPLITK_BLOCK_SIZE"] = 2 * kcfg["NUM_KSPLIT"] = 1cfg["BLOCK_SIZE_N"] = max(cfg["BLOCK_SIZE_N"], 32)if cfg["NUM_KSPLIT"] <= 1:cfg["NUM_KSPLIT"] = 1- cfg["SPLITK_BLOCK_SIZE"] = 2 * runtime_k+ cfg["SPLITK_BLOCK_SIZE"] = 2 * kreturn cfg- def _cfg_brief(cfg: dict[str, object]) -> str:- return (- f"bm={cfg['BLOCK_SIZE_M']},bn={cfg['BLOCK_SIZE_N']},bk={cfg['BLOCK_SIZE_K']},"- f"sp={cfg['NUM_KSPLIT']},sb={cfg['SPLITK_BLOCK_SIZE']},st={cfg['num_stages']},"- f"wp={cfg['num_warps']},wpe={cfg['waves_per_eu']}"- )+ # ── Pre-warming ─────────────────────────────────────────────────────────────-- def _emit_path(shape: tuple[int, int, int], path: str, cfg: dict[str, object], detail: str = "") -> None:- previous = _LOGGED_PATHS.get(shape)- if previous is not None:+ def _prewarm_all():+ """Nuclear pre-warming with selective disable-lsr."""+ if _DIRECT_KERNEL is None:return- _LOGGED_PATHS[shape] = path- suffix = f" {detail}" if detail else ""- print(- f"[amd2-mm-v244] shape={shape} path={path} {_cfg_brief(cfg)}{suffix}",- file=sys.stderr,- flush=True,- )+ all_m = [1, 2, 4, 8, 16, 32, 64, 128, 256]+ all_n = [2112, 2880, 3072, 4096, 7168]+ all_k = [512, 1536, 2048, 7168]- def _run_direct_kernel_path(- a_bf16: torch.Tensor,- b_shuffle: torch.Tensor,- b_scale_sh: torch.Tensor,- m: int,- n: int,- k: int,- ) -> tuple[torch.Tensor, dict[str, object], int, int]:- _resolve_runtime()+ phase1_cfgs = set() # no disable-lsr (M<=32 K>=1536)+ phase2_cfgs = set() # with disable-lsr (everything else)- shape = (m, n, k)- if not _DIRECT_KERNEL_SHAPE_SUPPORT.get(shape, True):- raise RuntimeError(f"direct kernel disabled for {shape}")+ 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)- if _DIRECT_KERNEL is None or _TRITON is None:- raise RuntimeError("direct Triton preshuffle kernel unavailable")+ # Phase 1: compile without disable-lsr+ prev = _set_disable_lsr(False)+ _prewarm_configs(phase1_cfgs)+ _restore_disable_lsr(prev)- b_ps_u8, s_ps_u8 = _get_cached_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- cfg = _prepare_direct_cfg(m, n, k, runtime_k)- if cfg["NUM_KSPLIT"] > 1 and _REDUCE_KERNEL is None:- raise RuntimeError("direct Triton reduce kernel unavailable")- y = _get_cached_output(a_bf16.device, m, runtime_n)+ # Phase 2: compile with disable-lsr+ prev = _set_disable_lsr(True)+ _prewarm_configs(phase2_cfgs)+ _restore_disable_lsr(prev)- if cfg["NUM_KSPLIT"] > 1:- y_pp = _get_cached_partials(a_bf16.device, int(cfg["NUM_KSPLIT"]), m, runtime_n)- out = y_pp- else:- y_pp = None- out = y+ # Pre-warm reduce kernel+ if _REDUCE_KERNEL is not None:+ _prewarm_reduce()- stride_am = k- stride_ak = 1- stride_bn = b_ps_u8.shape[1]- stride_bk = 1- stride_bsn = s_ps_u8.shape[1]- stride_bsk = 1- stride_cm = runtime_n- stride_cn = 1- if y_pp is None:- stride_ck = 0- launch_stride_cm = stride_cm- launch_stride_cn = stride_cn- else:- stride_ck = m * runtime_n- launch_stride_cm = runtime_n- launch_stride_cn = 1+ print(f"[mm-opt] Pre-warmed {len(phase1_cfgs)} no-lsr + {len(phase2_cfgs)} lsr configs",+ file=sys.stderr, flush=True)- block_size_m = cfg["BLOCK_SIZE_M"]- block_size_n = cfg["BLOCK_SIZE_N"]- block_size_k = cfg["BLOCK_SIZE_K"]- group_size_m = cfg["GROUP_SIZE_M"]- num_ksplit = cfg["NUM_KSPLIT"]- splitk_block_size = cfg["SPLITK_BLOCK_SIZE"]- num_stages = cfg["num_stages"]- num_warps = cfg["num_warps"]- waves_per_eu = cfg["waves_per_eu"]- matrix_instr_nonkdim = cfg["matrix_instr_nonkdim"]- cache_modifier = cfg["cache_modifier"]- grid = lambda meta: ( # noqa: E731- (- meta["NUM_KSPLIT"]- * _ceil_div(m, int(meta["BLOCK_SIZE_M"]))- * _ceil_div(runtime_n, int(meta["BLOCK_SIZE_N"]))- ),- )+ 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")- previous_disable_lsr = _set_disable_lsr(_shape_uses_disable_lsr(m, k))- try:- _DIRECT_KERNEL[grid](- a_bf16,- b_ps_u8,- out,- s_ps_u8,- m,- runtime_n,- runtime_k,- stride_am,- stride_ak,- stride_bn,- stride_bk,- stride_ck,- launch_stride_cm,- launch_stride_cn,- stride_bsn,- stride_bsk,- PREQUANT=True,- BLOCK_SIZE_M=block_size_m,- BLOCK_SIZE_N=block_size_n,- BLOCK_SIZE_K=block_size_k,- GROUP_SIZE_M=group_size_m,- NUM_KSPLIT=num_ksplit,- SPLITK_BLOCK_SIZE=splitk_block_size,- num_stages=num_stages,- num_warps=num_warps,- waves_per_eu=waves_per_eu,- matrix_instr_nonkdim=matrix_instr_nonkdim,- cache_modifier=cache_modifier,- )-- if y_pp is not None:- actual_ksplit = int(_TRITON.cdiv(runtime_k, int(cfg["SPLITK_BLOCK_SIZE"]) // 2))- grid_reduce = (_ceil_div(m, 16), _ceil_div(runtime_n, 16))- _REDUCE_KERNEL[grid_reduce](- y_pp,- y,- m,- runtime_n,- stride_ck,- launch_stride_cm,- launch_stride_cn,- stride_cm,- stride_cn,- 16,- 16,- actual_ksplit,- int(_TRITON.next_power_of_2(int(cfg["NUM_KSPLIT"]))),+ 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",)- return y, cfg, runtime_n, runtime_k- except Exception:- _DIRECT_KERNEL_SHAPE_SUPPORT[shape] = False- raise- finally:- _restore_disable_lsr(previous_disable_lsr)+ except Exception:+ pass- def _run_direct_helper_path(- a_bf16: torch.Tensor,- b_shuffle: torch.Tensor,- b_scale_sh: torch.Tensor,- m: int,- n: int,- k: int,- ) -> tuple[torch.Tensor, dict[str, object], int, int]:- _resolve_runtime()- if _DIRECT_HELPER is None:- raise RuntimeError("direct helper unavailable")+ 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- shape = (m, n, k)- if not _DIRECT_HELPER_SHAPE_SUPPORT.get(shape, True):- raise RuntimeError(f"direct helper disabled for {shape}")- cfg = _prepare_helper_cfg(m, n, k)- b_ps_u8, s_ps_u8 = _get_cached_preshuffle_views(b_shuffle, b_scale_sh, n, k)- y = _get_cached_output(a_bf16.device, m, n)+ # ── Fallback ────────────────────────────────────────────────────────────────- try:- config_arg = cfg- if not _DIRECT_HELPER_ACCEPTS_DICT and _SERIALIZE_DICT is not None:- config_arg = _SERIALIZE_DICT(cfg)- out = _DIRECT_HELPER(- a_bf16,- b_ps_u8,- s_ps_u8,- prequant=True,- dtype=torch.bfloat16,- y=y,- config=config_arg,- skip_reduce=False,- )- return out, cfg, n, (k // 2)- except Exception:- _DIRECT_HELPER_SHAPE_SUPPORT[shape] = False- raise+ 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_q: torch.Tensor,- b_shuffle: torch.Tensor,- a_scale_sh: torch.Tensor,- b_scale_sh: torch.Tensor,- m: int,- n: int,- k: int,- ) -> torch.Tensor:- return aiter.gemm_a4w4(- a_q,- b_shuffle,- a_scale_sh,- b_scale_sh,- dtype=dtypes.bf16,- bpreshuffle=True,- )+ 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⋯ 4 unchanged linesn = B_shuffle.shape[0]shape = (m, n, k)- try:- out, direct_cfg, runtime_n, runtime_k = _run_direct_kernel_path(A, B_shuffle, B_scale_sh, m, n, k)- _emit_path(shape, "direct", direct_cfg, f"rn={runtime_n},rk={runtime_k}")- return out- except Exception as direct_exc:- direct_detail = repr(direct_exc)+ _resolve_runtime()+ 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)++ 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++ 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"+ else:+ os.environ.pop("DISABLE_LLVM_OPT", None)++ 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"])),+ )+try:- out, helper_cfg, runtime_n, runtime_k = _run_direct_helper_path(A, B_shuffle, B_scale_sh, m, n, k)- _emit_path(shape, "helper", helper_cfg, f"rn={runtime_n},rk={runtime_k},direct={direct_detail}")- return out- except Exception as helper_exc:- helper_cfg = _prepare_helper_cfg(m, n, k)- A_q, A_scale_sh = _get_cached_a_quant(A)- _emit_path(- shape,- "fallback",- helper_cfg,- f"rn={n},rk={(k // 2)},direct={direct_detail},helper={repr(helper_exc)}",+ _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,)- return _run_fallback_gemm(A_q, B_shuffle, A_scale_sh, B_scale_sh, m, n, k)++ 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 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)
scrolls · 1079 diff lines total
Best evidence level for this revision: reported
JSON