Skip to content
KernelIndex
Search⌘K

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
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
9.08µs
#117 of 1143
2026-04-05

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.

fp4Optimized MXFP4 GEMM submission based on GEMM-Reference.md best practices.
split-k_GET_SPLITK = None
tile-k = 2565. Tuned per-shape configs (BK=256 pipeline, KSPLIT routing, BM=8 for M<=32)
tile-m = 85. 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 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
- _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 lines
return 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: True
except 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 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:
- 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 = 4
elif 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 = 3
else:
ksplit = 1
- elif k >= 2048 and tiles_for_split > _CU and tiles_for_split <= (_CU * 3) // 2:
- ksplit = 2
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:
⋯ 1 unchanged lines
else:
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 selection
block_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_size
cfg["BLOCK_SIZE_K"] = block_size_k
cfg["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 * 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 * runtime_k
+ cfg["SPLITK_BLOCK_SIZE"] = 2 * k
return 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 lines
n = 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