Skip to content
KernelIndex
Search⌘K

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
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
11.4µs
#329 of 1143
2026-03-24

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.

autotunesubmission_6: self-contained fused A16 W-FP4 GEMM + optional in-code autotune.
fp4submission_6: self-contained fused A16 W-FP4 GEMM + optional in-code autotune.
split-kand (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)

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