submission 622717
Stas Polonsky · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 729 lines, June 9 Researcher Reciprocity License v1.0.
submission_6.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-622717?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:6dcd9a2369308a13a4bb982466851c0e2de5b03a93689349f3baa226d95611c3
license declaredunknown
license concludedunknown
authorsStas Polonsky
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
submission_6.py729 lines
"""
submission_6: self-contained fused A16 W-FP4 GEMM + optional in-code autotune.
"""
from __future__ import annotations
from typing import Optional
import time
import sys
import random
import torch
import triton
import triton.language as tl
from task import input_t, output_t
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
_gemm_afp4wfp4_reduce_kernel,
)
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid
from aiter.ops.triton.utils.gemm_config_utils import get_gemm_config
try:
from utils import clear_l2_cache_large as _clear_l2_cache
except Exception:
def _clear_l2_cache() -> None:
return None
# --- In-code tuning controls (no env vars) ---
_AUTO_TUNE_ENABLED = False
_AUTO_TUNE_DISABLE_AFTER_TUNE = True
_AUTO_TUNE_MAX_CANDIDATES = 60
_AUTO_TUNE_WARMUP = 1
_AUTO_TUNE_ITERS = 12
# Internal key: (M, N_logical, K_packed)
_AUTO_TUNE_TARGET_SHAPE: Optional[tuple[int, int, int]] = (256, 3072, 768)
_AUTO_TUNE_VERBOSE = True
# Baseline medians in microseconds (update these when you re-baseline).
# Mapping uses internal keys: (M, N_logical, K_packed) where K_packed = K_bf16 // 2.
_BASELINE_TIME_US_BY_SHAPE: dict[tuple[int, int, int], float] = {
(4, 2880, 256): 7.16, # k_bf16=512, m=4, n=2880 (best tuned)
(16, 2112, 3584): 14.40, # k_bf16=7168, m=16, n=2112 (best tuned)
(32, 4096, 256): 7.12, # k_bf16=512, m=32, n=4096 (best tuned)
(32, 2880, 256): 7.10, # k_bf16=512, m=32, n=2880 (best tuned)
(64, 7168, 1024): 17.58, # k_bf16=2048, m=64, n=7168 (best tuned)
(256, 3072, 768): 15.24, # k_bf16=1536, m=256, n=3072 (best tuned)
}
def _to_u8(t: torch.Tensor) -> torch.Tensor:
if t.dtype == torch.uint8:
return t
return t.view(torch.uint8)
def _prepare_b_for_fused(b: torch.Tensor, b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor):
n, k = b.shape
if (n % 32) != 0 or (k % 32) != 0:
raise RuntimeError(
f"fused_contract_violation: require N%32==0 and K%32==0, got N={n},K={k}"
)
b_shuffle_u8 = _to_u8(b_shuffle)
b_scale_sh_u8 = _to_u8(b_scale_sh)
w_shuffled = b_shuffle_u8.view(n // 16, (k // 2) * 16)
scales_rows = b_scale_sh_u8.shape[0] // 32
w_scales_full = b_scale_sh_u8.view(scales_rows, -1)
need_rows = n // 32
if w_scales_full.shape[0] < need_rows or w_scales_full.shape[1] < k:
raise RuntimeError(
"fused_contract_violation: insufficient B_scale_sh after view "
f"got={tuple(w_scales_full.shape)} need=({need_rows},{k})"
)
w_scales_shuffled = w_scales_full[:need_rows, :k]
return w_shuffled, w_scales_shuffled
@triton.heuristics(
{
"EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)
and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),
"GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])
* triton.cdiv(args["N"], args["BLOCK_SIZE_N"]),
}
)
@triton.jit
def _gemm_a16wfp4_preshuffle_kernel(
a_ptr,
b_ptr,
c_ptr,
b_scales_ptr,
M,
N,
K,
stride_am,
stride_ak,
stride_bn,
stride_bk,
stride_ck,
stride_cm,
stride_cn,
stride_bsn,
stride_bsk,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr,
NUM_KSPLIT: tl.constexpr,
SPLITK_BLOCK_SIZE: tl.constexpr,
EVEN_K: tl.constexpr,
num_warps: tl.constexpr,
num_stages: tl.constexpr,
waves_per_eu: tl.constexpr,
matrix_instr_nonkdim: tl.constexpr,
GRID_MN: tl.constexpr,
PREQUANT: tl.constexpr,
cache_modifier: tl.constexpr,
):
tl.assume(stride_am > 0)
tl.assume(stride_ak > 0)
tl.assume(stride_bk > 0)
tl.assume(stride_bn > 0)
tl.assume(stride_cm > 0)
tl.assume(stride_cn > 0)
tl.assume(stride_bsk > 0)
tl.assume(stride_bsn > 0)
pid_unified = tl.program_id(axis=0)
pid_k = pid_unified % NUM_KSPLIT
pid = pid_unified // NUM_KSPLIT
num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
if NUM_KSPLIT == 1:
pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
else:
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
SCALE_GROUP_SIZE: tl.constexpr = 32
if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16
offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
a_ptrs = a_ptr + (
offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak
)
offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
b_ptrs = b_ptr + (
offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk
)
offs_bsn = (
pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, (BLOCK_SIZE_N // 32))
) % N
offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(
0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32
)
b_scale_ptrs = (
b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk
)
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for k_iter in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
b_scales = (
tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
.reshape(
BLOCK_SIZE_N // 32,
BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8,
4,
16,
2,
2,
1,
)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
)
if EVEN_K:
a_bf16 = tl.load(a_ptrs)
b_vals = tl.load(b_ptrs, cache_modifier=cache_modifier)
else:
a_bf16 = tl.load(
a_ptrs,
mask=offs_k_bf16[None, :] < 2 * K - k_iter * BLOCK_SIZE_K,
other=0,
)
b_vals = tl.load(
b_ptrs,
mask=offs_k_shuffle_arr[None, :] < (K - k_iter * (BLOCK_SIZE_K // 2)) * 16,
other=0,
cache_modifier=cache_modifier,
)
b_vals = (
b_vals.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
.trans(1, 0)
)
if PREQUANT:
a_q, a_scales = _mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
accumulator += tl.dot_scaled(a_q, a_scales, "e2m1", b_vals, b_scales, "e2m1")
a_ptrs += BLOCK_SIZE_K * stride_ak
b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
b_scale_ptrs += BLOCK_SIZE_K * stride_bsk
c = accumulator.to(c_ptr.type.element_ty)
offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
c_ptrs = (
c_ptr
+ stride_cm * offs_cm[:, None]
+ stride_cn * offs_cn[None, :]
+ pid_k * stride_ck
)
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
tl.store(c_ptrs, c, mask=c_mask)
def _get_splitk(K: int, block_size_k: int, num_ksplit: int):
splitk_block_size = (
triton.cdiv((2 * triton.cdiv(K, num_ksplit)), block_size_k) * block_size_k
)
while num_ksplit > 1 and block_size_k > 16:
if (
K % (splitk_block_size // 2) == 0
and splitk_block_size % block_size_k == 0
and K % (block_size_k // 2) == 0
):
break
if K % (splitk_block_size // 2) != 0 and num_ksplit > 1:
num_ksplit = num_ksplit // 2
elif splitk_block_size % block_size_k != 0:
if num_ksplit > 1:
num_ksplit = num_ksplit // 2
elif block_size_k > 16:
block_size_k = block_size_k // 2
elif K % (block_size_k // 2) != 0 and block_size_k > 16:
block_size_k = block_size_k // 2
else:
break
splitk_block_size = (
triton.cdiv((2 * triton.cdiv(K, num_ksplit)), block_size_k) * block_size_k
)
return splitk_block_size, block_size_k, num_ksplit
def _get_config(M: int, N: int, K: int):
return get_gemm_config("GEMM-A16WFP4_PRESHUFFLED", M, N, 2 * K)
# Test-shape overrides in evaluation order.
# Test 1
_SHAPE_4_2880_256_CONFIG = {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 4,
}
# Test 2
_SHAPE_16_2112_3584_CONFIG = {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 8,
}
# Test 3
_SHAPE_32_4096_256_CONFIG = {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
}
# Test 4
_SHAPE_32_2880_256_CONFIG = {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 1,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": None,
"NUM_KSPLIT": 4,
}
# Test 5
_SHAPE_64_7168_1024_CONFIG = {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 2,
"waves_per_eu": 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
}
# Test 6
_SHAPE_256_3072_768_CONFIG = {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 512,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 4,
"waves_per_eu": 2,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
"NUM_KSPLIT": 1,
}
_SHAPE_CONFIG_OVERRIDES: dict[tuple[int, int, int], dict] = {
(4, 2880, 256): _SHAPE_4_2880_256_CONFIG,
(16, 2112, 3584): _SHAPE_16_2112_3584_CONFIG,
(32, 4096, 256): _SHAPE_32_4096_256_CONFIG,
(32, 2880, 256): _SHAPE_32_2880_256_CONFIG,
(64, 7168, 1024): _SHAPE_64_7168_1024_CONFIG,
(256, 3072, 768): _SHAPE_256_3072_768_CONFIG,
}
def gemm_a16wfp4_preshuffle_local(
x: torch.Tensor,
w: torch.Tensor,
w_scales: torch.Tensor,
prequant: Optional[bool] = True,
dtype: Optional[torch.dtype] = torch.bfloat16,
y: Optional[torch.Tensor] = None,
config: Optional[dict] = None,
skip_reduce: Optional[bool] = False,
) -> torch.Tensor:
assert prequant, "prequant == False is not supported"
M, K = x.shape
N, K = w.shape
N = N * 16
K = K // 16
if config is None:
key = (M, N, K)
if key in _SHAPE_CONFIG_OVERRIDES:
config = dict(_SHAPE_CONFIG_OVERRIDES[key])
else:
config, _ = _get_config(M, N, K)
if config["NUM_KSPLIT"] > 1:
splitk_block_size, block_size_k, num_ksplit = _get_splitk(
K, config["BLOCK_SIZE_K"], config["NUM_KSPLIT"]
)
config["SPLITK_BLOCK_SIZE"] = splitk_block_size
config["BLOCK_SIZE_K"] = block_size_k
config["NUM_KSPLIT"] = num_ksplit
if config["BLOCK_SIZE_K"] >= 2 * K:
config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * K)
config["SPLITK_BLOCK_SIZE"] = 2 * K
config["NUM_KSPLIT"] = 1
config["BLOCK_SIZE_N"] = max(config["BLOCK_SIZE_N"], 32)
return_y_pp = config["NUM_KSPLIT"] > 1 and skip_reduce
if config["NUM_KSPLIT"] > 1:
y_pp = torch.empty((config["NUM_KSPLIT"], M, N), dtype=torch.float32, device=x.device)
else:
config["SPLITK_BLOCK_SIZE"] = 2 * K
y_pp = None
if y is None and not return_y_pp:
y = torch.empty((M, N), dtype=dtype, device=x.device)
grid = lambda META: ( # noqa: E731
META["NUM_KSPLIT"]
* triton.cdiv(M, META["BLOCK_SIZE_M"])
* triton.cdiv(N, META["BLOCK_SIZE_N"]),
)
_gemm_a16wfp4_preshuffle_kernel[grid](
x,
w,
y if y_pp is None else y_pp,
w_scales,
M,
N,
K,
x.stride(0),
x.stride(1),
w.stride(0),
w.stride(1),
0 if y_pp is None else y_pp.stride(0),
y.stride(0) if y_pp is None else y_pp.stride(1),
y.stride(1) if y_pp is None else y_pp.stride(2),
w_scales.stride(0),
w_scales.stride(1),
PREQUANT=prequant,
**config,
)
if return_y_pp:
return y_pp
if config["NUM_KSPLIT"] > 1:
reduce_block_size_m = 16
reduce_block_size_n = 64
actual_ksplit = triton.cdiv(K, (config["SPLITK_BLOCK_SIZE"] // 2))
grid_reduce = (
triton.cdiv(M, reduce_block_size_m),
triton.cdiv(N, reduce_block_size_n),
)
_gemm_afp4wfp4_reduce_kernel[grid_reduce](
y_pp,
y,
M,
N,
y_pp.stride(0),
y_pp.stride(1),
y_pp.stride(2),
y.stride(0),
y.stride(1),
reduce_block_size_m,
reduce_block_size_n,
actual_ksplit,
triton.next_power_of_2(config["NUM_KSPLIT"]),
)
return y
# --- In-code result recording ---
_TUNE_HISTORY: dict[tuple[int, int, int], list[dict]] = {}
_BEST_CONFIG_BY_SHAPE: dict[tuple[int, int, int], dict] = {}
_BEST_TIME_US_BY_SHAPE: dict[tuple[int, int, int], float] = {}
_TUNED_SHAPES: set[tuple[int, int, int]] = set()
_LAST_TUNE_SEED_BY_SHAPE: dict[tuple[int, int, int], int] = {}
def _shape_matches_target(key: tuple[int, int, int]) -> bool:
return _AUTO_TUNE_TARGET_SHAPE is None or key == _AUTO_TUNE_TARGET_SHAPE
def _median_us(samples_us: list[float]) -> float:
values = sorted(samples_us)
n = len(values)
if n == 0:
return float("inf")
mid = n // 2
if n % 2 == 1:
return values[mid]
return 0.5 * (values[mid - 1] + values[mid])
def _default_config_for_key(key: tuple[int, int, int]) -> dict:
m, n, k = key
cfg, _ = _get_config(m, n, k)
return dict(cfg)
def _dedupe_candidates(candidates: list[dict]) -> list[dict]:
out: list[dict] = []
seen: set[tuple] = set()
for cfg in candidates:
sig = tuple(sorted(cfg.items()))
if sig in seen:
continue
seen.add(sig)
out.append(cfg)
return out
def _build_candidates(key: tuple[int, int, int]) -> list[dict]:
base = dict(_SHAPE_CONFIG_OVERRIDES.get(key, _default_config_for_key(key)))
candidates: list[dict] = [dict(base)]
knobs = {
"BLOCK_SIZE_N": [64, 128, 256],
"BLOCK_SIZE_M": [16, 32, 64],
"num_stages": [1, 2, 4],
"NUM_KSPLIT": [1, 2, 4, 8, 14],
"cache_modifier": [None, ".cg"],
"num_warps": [4, 8],
"waves_per_eu": [1, 2, 4],
}
for k, values in knobs.items():
for v in values:
if base.get(k) == v:
continue
cfg = dict(base)
cfg[k] = v
candidates.append(cfg)
deduped = _dedupe_candidates(candidates)
# Re-shuffle candidates on each tuning call so we can explore different subsets
# when MAX_CANDIDATES truncates the search space.
seed = time.time_ns() ^ (key[0] * 1_000_003) ^ (key[1] * 10_007) ^ key[2]
rng = random.Random(seed)
rng.shuffle(deduped)
_LAST_TUNE_SEED_BY_SHAPE[key] = seed
return deduped[: _AUTO_TUNE_MAX_CANDIDATES]
def _format_best_configs_compact() -> str:
if not _BEST_CONFIG_BY_SHAPE:
return "_SHAPE_CONFIG_OVERRIDES: dict[tuple[int, int, int], dict] = {}"
lines = ["_SHAPE_CONFIG_OVERRIDES: dict[tuple[int, int, int], dict] = {"]
for key in sorted(_BEST_CONFIG_BY_SHAPE):
cfg = _BEST_CONFIG_BY_SHAPE[key]
best_us = _BEST_TIME_US_BY_SHAPE.get(key, float("nan"))
lines.append(f" {key}: {cfg}, # best_us={best_us:.2f}")
lines.append("}")
return "\n".join(lines)
def print_best_configs_compact() -> None:
"""Print copy-paste snippet for freezing tuned configs into static overrides."""
print("[autotune-freeze] begin", file=sys.stderr, flush=True)
print(_format_best_configs_compact(), file=sys.stderr, flush=True)
print("[autotune-freeze] end", file=sys.stderr, flush=True)
def _print_tune_history_compact(key: tuple[int, int, int]) -> None:
history = _TUNE_HISTORY.get(key, [])
if not history:
return
ok_rows = [r for r in history if r.get("status") == "ok"]
err_rows = [r for r in history if r.get("status") != "ok"]
ok_rows.sort(key=lambda r: float(r.get("time_us", float("inf"))))
print(
f"[autotune-history] shape={key} tried={len(history)} ok={len(ok_rows)} err={len(err_rows)}",
file=sys.stderr,
flush=True,
)
for rank, row in enumerate(ok_rows, start=1):
print(
f"[autotune-history] rank={rank} idx={row['idx']} time_us={row['time_us']:.2f} config={row['config']}",
file=sys.stderr,
flush=True,
)
for row in err_rows:
print(
f"[autotune-history] idx={row['idx']} status=error error={row.get('error')} config={row['config']}",
file=sys.stderr,
flush=True,
)
def _time_candidate(
a: torch.Tensor,
w_shuffled: torch.Tensor,
w_scales_shuffled: torch.Tensor,
out_dtype: torch.dtype,
cfg: dict,
) -> float:
# Warmup/JIT
for _ in range(_AUTO_TUNE_WARMUP):
_ = gemm_a16wfp4_preshuffle_local(
a,
w_shuffled,
w_scales_shuffled,
prequant=True,
dtype=out_dtype,
config=dict(cfg),
)
torch.cuda.synchronize()
samples_us: list[float] = []
for _ in range(_AUTO_TUNE_ITERS):
torch.cuda.synchronize()
_clear_l2_cache()
start_event = torch.cuda.Event(enable_timing=True)
end_event = torch.cuda.Event(enable_timing=True)
start_event.record()
_ = gemm_a16wfp4_preshuffle_local(
a,
w_shuffled,
w_scales_shuffled,
prequant=True,
dtype=out_dtype,
config=dict(cfg),
)
end_event.record()
torch.cuda.synchronize()
samples_us.append(start_event.elapsed_time(end_event) * 1000.0)
return _median_us(samples_us)
def _auto_tune_shape(
key: tuple[int, int, int],
a: torch.Tensor,
w_shuffled: torch.Tensor,
w_scales_shuffled: torch.Tensor,
out_dtype: torch.dtype,
) -> None:
candidates = _build_candidates(key)
history: list[dict] = []
best_time = float("inf")
best_cfg: Optional[dict] = None
for idx, cfg in enumerate(candidates):
try:
median_us = _time_candidate(a, w_shuffled, w_scales_shuffled, out_dtype, cfg)
rec = {
"idx": idx,
"status": "ok",
"time_us": median_us,
"config": dict(cfg),
}
if median_us < best_time:
best_time = median_us
best_cfg = dict(cfg)
except Exception as exc:
rec = {
"idx": idx,
"status": "error",
"error": f"{type(exc).__name__}: {exc}",
"config": dict(cfg),
}
history.append(rec)
_TUNE_HISTORY[key] = history
_TUNED_SHAPES.add(key)
if best_cfg is not None:
baseline_us = _BASELINE_TIME_US_BY_SHAPE.get(key)
is_better_than_baseline = baseline_us is None or best_time < baseline_us
if is_better_than_baseline:
_BEST_CONFIG_BY_SHAPE[key] = best_cfg
_BEST_TIME_US_BY_SHAPE[key] = best_time
if _AUTO_TUNE_VERBOSE:
print(
f"[autotune] shape={key} seed={_LAST_TUNE_SEED_BY_SHAPE.get(key)} "
f"best_us={best_time:.2f} config={best_cfg}",
file=sys.stderr,
flush=True,
)
if baseline_us is not None:
delta = best_time - baseline_us
if is_better_than_baseline:
print(
f"[autotune-compare] shape={key} baseline_us={baseline_us:.2f} "
f"best_us={best_time:.2f} delta_us={delta:.2f} decision=update",
file=sys.stderr,
flush=True,
)
else:
print(
f"[autotune-compare] shape={key} baseline_us={baseline_us:.2f} "
f"best_us={best_time:.2f} delta_us={delta:.2f} decision=keep_baseline",
file=sys.stderr,
flush=True,
)
_print_tune_history_compact(key)
# Print compact final-best table (one best config per shape) for copy/paste.
print_best_configs_compact()
elif _AUTO_TUNE_VERBOSE:
print(
f"[autotune] shape={key} no valid candidate, fallback to static/default config",
file=sys.stderr,
flush=True,
)
def custom_kernel(data: input_t) -> output_t:
global _AUTO_TUNE_ENABLED
from aiter import dtypes
a, b, _b_q, b_shuffle, b_scale_sh = data
a = a.contiguous()
b = b.contiguous()
b_shuffle = b_shuffle.contiguous()
b_scale_sh = b_scale_sh.contiguous()
w_shuffled, w_scales_shuffled = _prepare_b_for_fused(b, b_shuffle, b_scale_sh)
# Internal key format used by shape overrides.
key = (a.shape[0], b.shape[0], b.shape[1] // 2)
if _AUTO_TUNE_ENABLED and _shape_matches_target(key) and key not in _TUNED_SHAPES:
_auto_tune_shape(
key=key,
a=a,
w_shuffled=w_shuffled,
w_scales_shuffled=w_scales_shuffled,
out_dtype=dtypes.bf16,
)
if _AUTO_TUNE_DISABLE_AFTER_TUNE:
_AUTO_TUNE_ENABLED = False
# Priority: tuned best -> static override -> default loader.
config = None
if key in _BEST_CONFIG_BY_SHAPE:
config = dict(_BEST_CONFIG_BY_SHAPE[key])
elif key in _SHAPE_CONFIG_OVERRIDES:
config = dict(_SHAPE_CONFIG_OVERRIDES[key])
return gemm_a16wfp4_preshuffle_local(
a,
w_shuffled,
w_scales_shuffled,
prequant=True,
dtype=dtypes.bf16,
config=config,
)
scrolls · 729 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON