Skip to content
KernelIndex
Search⌘K

submission 754371

.jonnss · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 537 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-754371?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, mxfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
8.81µs
#88 of 1143
2026-04-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a8877a15a922f0ce34cbcec5af4546266b27407cbdb39a4c62de2ac2d86b1632
license declaredunknown
license concludedunknown
authors.jonnss
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fp4CDNA4 blocked FP4 matmul via aiter preshuffle path.
num-warps = 4num_warps=4, num_stages=2, waves_per_eu=wpe,
split-k_rows = ["cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio"]
stages = 2num_warps=4, num_stages=2, waves_per_eu=wpe,
vector-width = float4float4 acc = *reinterpret_cast<const float4*>(src + base);

Kernel source

submission.py537 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
CDNA4 blocked FP4 matmul via aiter preshuffle path.
Runtime patches the quantization subroutine with ISA-level
paired conversion and adjusts the accumulator idiom.
Occupancy-aware tiling with phased JIT warmup.
"""

import os as _env
_env.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
_env.environ.setdefault("CXX", "clang++")
import uuid as _uid
_env.environ["TRITON_CACHE_DIR"] = f"/tmp/_tc_{_uid.uuid4().hex[:8]}"

# Synthesize a minimal config CSV so aiter skips its heavy build paths
_ASM_LABEL = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_CSV_TMP = "/tmp/_fp4_cfg.csv"
_ENGINES = 256
_DIM_PAIRS = [
    (2880, 512), (2112, 7168), (4096, 512), (7168, 2048), (3072, 1536),
    (2880, 1536), (4096, 1536), (2112, 512), (2112, 2048),
    (7168, 512), (7168, 1536), (7168, 7168), (3072, 512),
    (3072, 7168), (3072, 2048), (4096, 2048), (4096, 7168),
    (2880, 2048), (2880, 7168),
]
_BATCH_DIMS = [1, 2, 4, 8, 16, 32, 64, 128, 256]
_rows = ["cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio"]
for _d1, _d2 in _DIM_PAIRS:
    for _b in _BATCH_DIMS:
        _wave_cnt = ((_b + 31) // 32) * ((_d1 + 127) // 128)
        _ratio = _ENGINES / max(_wave_cnt, 1)
        _lg = 0
        while _ratio >= pow(2, _lg + 1) and (pow(2, _lg + 1) * 128) < 2 * _d2:
            _lg += 1
        _lg = min(_lg, 3)
        _rows.append(f"{_ENGINES},{_b},{_d1},{_d2},21,{_lg},1.0,{_ASM_LABEL},0,0,0.0")
with open(_CSV_TMP, "w") as _fh:
    _fh.write("\n".join(_rows))
_env.environ["AITER_CONFIG_GEMM_A4W4"] = (
    _CSV_TMP + ":/home/runner/aiter/aiter/configs/a4w4_blockscale_tuned_gemm.csv"
)

import torch
torch.set_grad_enabled(False)
import triton
import triton.language as tl
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (
    _gemm_a16wfp4_preshuffle_kernel,
)
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
    _gemm_afp4wfp4_reduce_kernel,
)
from task import input_t, output_t
import sys as _io
import time as _clk
import gc as _mem
_io.setswitchinterval(1.0)
_log = lambda s: print(s, file=_io.stderr, flush=True)


# ---- Override heuristics to avoid dynamic tile selection ----
try:
    _gemm_a16wfp4_preshuffle_kernel.values['GRID_MN'] = lambda args: 1
    _gemm_a16wfp4_preshuffle_kernel.values['EVEN_K'] = lambda args: True
    _log("[setup] heuristics locked")
except Exception as _exc:
    _log(f"[setup] heuristic override failed: {_exc}")

_env.environ["HIP_FORCE_DEV_KERNARG"] = "1"


# ---- Inject ISA-level BF16->FP4 quantization ----
_log("[setup] patching quantizer with hardware conversion...")
try:
    _kern_obj = (
        _gemm_a16wfp4_preshuffle_kernel.fn
        if hasattr(_gemm_a16wfp4_preshuffle_kernel, 'fn')
        else _gemm_a16wfp4_preshuffle_kernel
    )
    _quant_ref = _kern_obj.__globals__['_mxfp4_quant_op']

    _patched_body = '''def _mxfp4_quant_op(
    x,
    BLOCK_SIZE_N,
    BLOCK_SIZE_M,
    MXFP4_QUANT_BLOCK_SIZE,
):
    """ISA-accelerated BF16 to packed FP4 via v_cvt_scalef32_pk_fp4_bf16."""
    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
    HALF_BLOCK: tl.constexpr = MXFP4_QUANT_BLOCK_SIZE // 2

    x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)

    amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
    amax = amax.to(tl.int32, bitcast=True)
    amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000

    amax_exp = (amax >> 23) & 0xFF
    scale_e8m0_unbiased = (amax_exp.to(tl.int32) - 129).to(tl.float32)
    scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)

    bs_e8m0 = (scale_e8m0_unbiased + 127).to(tl.float32).to(tl.uint8)

    biased_exp_f = tl.maximum(scale_e8m0_unbiased + 127.0, 1.0)
    hw_scale = (biased_exp_f.to(tl.int32).to(tl.uint32) << 23).to(tl.float32, bitcast=True)

    x_bf16 = x.to(tl.bfloat16)
    x_pairs = x_bf16.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, HALF_BLOCK, 2)
    evens, odds = tl.split(x_pairs)
    lo = evens.to(tl.uint16, bitcast=True).to(tl.uint32)
    hi = odds.to(tl.uint16, bitcast=True).to(tl.uint32)
    packed_bf16 = lo | (hi << 16)

    result = tl.inline_asm_elementwise(
        "v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2",
        "=v,v,v",
        [packed_bf16, hw_scale],
        dtype=tl.uint32,
        is_pure=True,
        pack=1,
    )

    x_fp4 = (result & 0xFF).to(tl.uint8)
    x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)

    return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
'''

    if hasattr(_quant_ref, '_unsafe_update_src'):
        _quant_ref._unsafe_update_src(_patched_body)
    else:
        _quant_ref._src = _patched_body
        if hasattr(_quant_ref, 'src'):
            _quant_ref.src = _patched_body
        if hasattr(_quant_ref, 'hash'):
            _quant_ref.hash = None

    _orig_src = _kern_obj._src
    _mod_src = _orig_src.replace(
        'accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")',
        'accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", acc=accumulator)'
    )
    if _mod_src != _orig_src:
        _kern_obj._unsafe_update_src(_mod_src)
        _log("[setup] quantizer + kernel patched OK")
    else:
        _log("[setup] quantizer patched, kernel string mismatch")

    _check = _quant_ref._src if hasattr(_quant_ref, '_src') else ''
    _log(f"[setup] has hw asm: {'inline_asm_elementwise' in _check}")
except Exception as _exc:
    import traceback
    _log(f"[setup] patch FAILED: {_exc}")
    traceback.print_exc(file=_io.stderr)


# ---- Native HIP accumulator merger for K-partitioned runs ----
_MERGER_HIP = r"""
#include <hip/hip_runtime.h>

__device__ __forceinline__ unsigned short to_bf16(float val) {
    unsigned int raw;
    __builtin_memcpy(&raw, &val, sizeof(raw));
    unsigned int bias = ((raw >> 16) & 1) + 0x7FFFu;
    return (unsigned short)((raw + bias) >> 16);
}

template <int NP>
__global__ void vec4_sum(const float* __restrict__ src,
                         unsigned short* __restrict__ dst, int len) {
    int base = (blockIdx.x * blockDim.x + threadIdx.x) * 4;
    if (base + 3 < len) {
        float4 acc = *reinterpret_cast<const float4*>(src + base);
        #pragma unroll
        for (int p = 1; p < NP; p++) {
            float4 part = *reinterpret_cast<const float4*>(src + p * len + base);
            acc.x += part.x; acc.y += part.y; acc.z += part.z; acc.w += part.w;
        }
        unsigned short a = to_bf16(acc.x), b = to_bf16(acc.y);
        unsigned short c = to_bf16(acc.z), d = to_bf16(acc.w);
        *reinterpret_cast<unsigned long long*>(dst + base) =
            (unsigned long long)a | ((unsigned long long)b << 16) |
            ((unsigned long long)c << 32) | ((unsigned long long)d << 48);
    } else {
        for (int j = base; j < len && j < base + 4; j++) {
            float acc = src[j];
            #pragma unroll
            for (int p = 1; p < NP; p++) acc += src[p * len + j];
            dst[j] = to_bf16(acc);
        }
    }
}

__global__ void scalar_sum(const float* __restrict__ src,
                           unsigned short* __restrict__ dst,
                           int len, int np) {
    int gid = blockIdx.x * blockDim.x + threadIdx.x;
    if (gid < len) {
        float acc = src[gid];
        for (int p = 1; p < np; p++) acc += src[p * len + gid];
        dst[gid] = to_bf16(acc);
    }
}

void merge_partials(torch::Tensor src, torch::Tensor dst, int R, int C, int np) {
    int len = R * C;
    const float* sp = src.data_ptr<float>();
    unsigned short* dp = reinterpret_cast<unsigned short*>(dst.data_ptr());
    const int thr = 64, stride = thr * 4;
    const int nblk = (len + stride - 1) / stride;
    switch (np) {
        case 2: vec4_sum<2><<<nblk, thr>>>(sp, dp, len); break;
        case 3: vec4_sum<3><<<nblk, thr>>>(sp, dp, len); break;
        case 4: vec4_sum<4><<<nblk, thr>>>(sp, dp, len); break;
        case 7: vec4_sum<7><<<nblk, thr>>>(sp, dp, len); break;
        case 8: vec4_sum<8><<<nblk, thr>>>(sp, dp, len); break;
        default: {
            const int t2 = 256, b2 = (len + t2 - 1) / t2;
            scalar_sum<<<b2, t2>>>(sp, dp, len, np);
            break;
        }
    }
}
"""
_MERGER_HDR = "void merge_partials(torch::Tensor src, torch::Tensor dst, int R, int C, int np);"

_HAS_HIP_MERGER = False
try:
    from torch.utils.cpp_extension import load_inline as _jit
    _jit_t0 = _clk.time()
    _hip_merger = _jit(
        name="fp4_kmerge",
        cpp_sources=[_MERGER_HDR],
        cuda_sources=[_MERGER_HIP],
        functions=["merge_partials"],
        verbose=False,
        extra_cuda_cflags=["--offload-arch=gfx950", "-O3"],
    )
    _HAS_HIP_MERGER = True
    _log(f"[setup] HIP merger ready ({_clk.time()-_jit_t0:.1f}s)")
except Exception as _exc:
    _log(f"[setup] HIP merger unavailable: {_exc}")


# ---- Partition alignment utility ----
def _align_parts(kh, bk, np):
    span = triton.cdiv((2 * triton.cdiv(kh, np)), bk) * bk
    while np > 1 and bk > 16:
        ok = (kh % (span // 2) == 0 and span % bk == 0 and kh % (bk // 2) == 0)
        if ok:
            break
        elif kh % (span // 2) != 0 and np > 1:
            np //= 2
        elif span % bk != 0:
            np = np // 2 if np > 1 else np
            if np <= 1 and bk > 16:
                bk //= 2
        elif kh % (bk // 2) != 0 and bk > 16:
            bk //= 2
        else:
            break
        span = triton.cdiv((2 * triton.cdiv(kh, np)), bk) * bk
    return span, bk, np


# ---- Occupancy-driven config resolver ----
_resolved = {}

def _shape_config(batch, cols, depth):
    tag = (batch, cols, depth)
    if tag in _resolved:
        return _resolved[tag]
    kh = depth // 2

    if batch <= 32:
        bm, bn = 8, 128
        wave_est = ((batch + bm - 1) // bm) * ((cols + 127) // 128)
        np = 1
        if depth >= 4096:
            np = 7
        elif depth >= 2048:
            np = 2 if (wave_est * 2 >= (_ENGINES * 3) // 4 and wave_est * 2 <= _ENGINES) else 4
        elif depth >= 1536:
            np = 2 if (wave_est * 2 >= (_ENGINES * 3) // 4 and wave_est * 2 <= _ENGINES) else 3
        bk = 256 if depth <= np * 512 or (np == 2 and depth <= np * 1024) else 512
        if wave_est * np < (_ENGINES * 3) // 4:
            bn = 64
        total_wg = ((batch + bm - 1) // bm) * ((cols + bn - 1) // bn) * np
        wpe = 2 if total_wg > _ENGINES else 1
    else:
        bm = 16
        if batch <= 128:
            est16 = ((batch + 15) // 16) * ((cols + 127) // 128)
            if est16 < (_ENGINES * 3) // 4:
                bm = 8
        wave_est = ((batch + bm - 1) // bm) * ((cols + 127) // 128)
        bn, np = 128, 1
        if _ENGINES // 2 <= wave_est <= _ENGINES and (depth >= 7168 or (depth >= 2048 and bm == 8)):
            np = 2
        elif wave_est < _ENGINES // 2 and depth > 512:
            if depth >= 4096:
                np = 2 if wave_est * 2 >= _ENGINES else 7
            elif depth >= 2048:
                np = 2
            elif depth >= 1536:
                np = 3
        bk = 256 if depth <= max(np * 4096, 2048) else 512
        if wave_est * np < (_ENGINES * 3) // 4:
            bn = 64
        total_wg = ((batch + bm - 1) // bm) * ((cols + bn - 1) // bn) * np
        wpe = 2 if total_wg > _ENGINES else 1

    params = {
        "BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": max(bn, 32), "BLOCK_SIZE_K": bk,
        "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
        "waves_per_eu": wpe, "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg", "NUM_KSPLIT": np,
    }

    if params["NUM_KSPLIT"] > 1:
        span, bk2, np2 = _align_parts(kh, params["BLOCK_SIZE_K"], params["NUM_KSPLIT"])
        params["SPLITK_BLOCK_SIZE"] = span
        params["BLOCK_SIZE_K"] = bk2
        params["NUM_KSPLIT"] = np2

    if params["BLOCK_SIZE_K"] >= 2 * kh:
        params["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * kh)
        params["SPLITK_BLOCK_SIZE"] = 2 * kh
        params["NUM_KSPLIT"] = 1
    params["BLOCK_SIZE_N"] = max(params["BLOCK_SIZE_N"], 32)

    if params["NUM_KSPLIT"] == 1:
        params["SPLITK_BLOCK_SIZE"] = 2 * kh

    real_np, padded_np = None, None
    if params["NUM_KSPLIT"] > 1:
        real_np = triton.cdiv(kh, params["SPLITK_BLOCK_SIZE"] // 2)
        padded_np = triton.next_power_of_2(params["NUM_KSPLIT"])

    m_tiles = triton.cdiv(batch, params["BLOCK_SIZE_M"])
    n_tiles = triton.cdiv(cols, params["BLOCK_SIZE_N"])
    launch_grid = (params["NUM_KSPLIT"] * m_tiles * n_tiles,)
    red_grid = None
    if params["NUM_KSPLIT"] > 1:
        red_grid = (triton.cdiv(batch, 16), triton.cdiv(cols, 16))

    bundle = (
        params, real_np, padded_np, launch_grid, red_grid,
        kh, params["BLOCK_SIZE_M"], params["BLOCK_SIZE_N"],
        params["BLOCK_SIZE_K"], params["NUM_KSPLIT"],
        params["SPLITK_BLOCK_SIZE"], params["waves_per_eu"],
    )
    _resolved[tag] = bundle
    return bundle


# ---- Phased JIT warmup ----
_warm_t0 = _clk.time()
_warmed = {}
_phase1 = {}
_phase2 = {}
_red_set = set()

for _dn, _dk in _DIM_PAIRS:
    for _db in _BATCH_DIMS:
        _cb, _rn, _rp, _, _, _, _, _, _, _, _, _ = _shape_config(_db, _dn, _dk)
        _sig = (
            _cb["BLOCK_SIZE_M"], _cb["BLOCK_SIZE_N"], _cb["BLOCK_SIZE_K"],
            _cb["NUM_KSPLIT"], _cb["SPLITK_BLOCK_SIZE"], _cb["waves_per_eu"],
        )
        if _db <= 32 and _dk >= 1536:
            _phase1.setdefault(_sig, True)
        else:
            _phase2.setdefault(_sig, True)
        if _rn is not None:
            _red_set.add((_rn, _rp))

for _s in _phase1:
    _phase2.pop(_s, None)

_log(f"[warm] {len(_phase1)} phase1 + {len(_phase2)} phase2, {len(_red_set)} reducers")

_dummy_x = torch.zeros(32, 8192, dtype=torch.bfloat16, device="cuda")
_dummy_w = torch.zeros(16, 65536, dtype=torch.uint8, device="cuda")
_dummy_s = torch.zeros(16, 65536, dtype=torch.uint8, device="cuda")
_dummy_pp = torch.zeros(16, 32, 256, dtype=torch.float32, device="cuda")
_dummy_y = torch.zeros(32, 256, dtype=torch.bfloat16, device="cuda")

def _fire_warmup(bm, bn, bk, ks, spk, wpe):
    cfg = {
        "BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": bn, "BLOCK_SIZE_K": bk,
        "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
        "waves_per_eu": wpe, "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg", "NUM_KSPLIT": ks, "SPLITK_BLOCK_SIZE": spk,
    }
    target = _dummy_pp if ks > 1 else _dummy_y
    _gemm_a16wfp4_preshuffle_kernel[(max(ks, 1),)](
        _dummy_x, _dummy_w, target, _dummy_s, bm, bn, spk // 2,
        _dummy_x.stride(0), _dummy_x.stride(1),
        _dummy_w.stride(0), _dummy_w.stride(1),
        0 if ks <= 1 else _dummy_pp.stride(0),
        _dummy_y.stride(0) if ks <= 1 else _dummy_pp.stride(1),
        _dummy_y.stride(1) if ks <= 1 else _dummy_pp.stride(2),
        _dummy_s.stride(0), _dummy_s.stride(1),
        PREQUANT=True, **cfg,
    )

_log("[warm] phase 1 (no lsr)...")
for _sig in sorted(_phase1):
    try:
        _fire_warmup(*_sig)
        _warmed[_sig] = 1
        _log(f"  {_sig[0]}x{_sig[1]}x{_sig[2]} ks={_sig[3]} ({_clk.time()-_warm_t0:.0f}s)")
    except Exception as _exc:
        _log(f"  {_sig}: ERR {_exc}")

_env.environ["DISABLE_LLVM_OPT"] = "disable-lsr"
_log(f"[warm] phase 2 (with lsr) @ {_clk.time()-_warm_t0:.0f}s")

for _idx, _sig in enumerate(sorted(_phase2)):
    if _clk.time() - _warm_t0 > 200:
        _log(f"  timeout, {len(_phase2) - _idx} skipped")
        break
    try:
        _fire_warmup(*_sig)
        _warmed[_sig] = 2
        _log(f"  {_sig[0]}x{_sig[1]}x{_sig[2]} ks={_sig[3]} ({_clk.time()-_warm_t0:.0f}s)")
    except Exception as _exc:
        _log(f"  {_sig}: ERR {_exc}")

_log(f"[warm] reducers...")
for _rn, _rp in sorted(_red_set):
    if _clk.time() - _warm_t0 > 230:
        _log("  timeout")
        break
    try:
        _gemm_afp4wfp4_reduce_kernel[(1, 1)](
            _dummy_pp, _dummy_y, 16, 16,
            _dummy_pp.stride(0), _dummy_pp.stride(1), _dummy_pp.stride(2),
            _dummy_y.stride(0), _dummy_y.stride(1), 16, 16, _rn, _rp,
        )
    except Exception:
        pass

del _dummy_x, _dummy_w, _dummy_s, _dummy_pp, _dummy_y, _fire_warmup
del _phase1, _phase2, _red_set
torch.cuda.empty_cache()
_log(f"[warm] done: {len(_warmed)} configs in {_clk.time()-_warm_t0:.0f}s")

_mem.disable()


# ---- Runtime dispatch state ----
_wt_cache = {}
_dest_buf = {}
_frag_buf = {}
_seen = set()


def _prepare_wt(data):
    addr = data[3].data_ptr()
    if addr not in _wt_cache:
        n_dim = data[3].shape[0]
        k_bytes = data[3].shape[1]
        sr, sc = data[4].shape
        n_grp = n_dim // 32
        w_view = data[3].view(torch.uint8).reshape(n_dim // 16, k_bytes * 16)
        s_view = data[4].view(torch.uint8).reshape(sr // 32, sc * 32)[:n_grp].contiguous()
        _wt_cache[addr] = (w_view, s_view, w_view.stride(0), s_view.stride(0))
    return _wt_cache[addr]


def custom_kernel(data: input_t) -> output_t:
    X = data[0]
    if not X.is_contiguous():
        X = X.contiguous()
    ndims = X.ndim
    X_flat = X if ndims == 2 else X.view(-1, X.shape[-1])
    batch = X_flat.shape[0]
    cols = data[3].shape[0]
    depth = data[3].shape[1] * 2

    (params, real_np, padded_np, launch_grid, red_grid,
     kh, bm, bn, bk, ks, spk, wpe) = _shape_config(batch, cols, depth)

    tag = (batch, cols, depth)
    if tag not in _seen:
        _seen.add(tag)
        _log(f"[run] {batch}x{cols}x{depth} bm={bm} bn={bn} bk={bk} ks={ks} wpe={wpe}")

    okey = (batch, cols)
    if okey not in _dest_buf:
        _dest_buf[okey] = torch.empty((batch, cols), dtype=torch.bfloat16, device="cuda")
    dest = _dest_buf[okey]

    w_view, s_view, sw0, ss0 = _prepare_wt(data)

    if ks > 1:
        fkey = (padded_np, batch, cols)
        if fkey not in _frag_buf:
            _frag_buf[fkey] = torch.empty(
                (padded_np, batch, cols), dtype=torch.float32, device="cuda"
            )
        frags = _frag_buf[fkey]
        stride_p, stride_r = batch * cols, cols
    else:
        frags = None
        stride_p, stride_r = 0, cols

    _gemm_a16wfp4_preshuffle_kernel[launch_grid](
        X_flat, w_view,
        dest if frags is None else frags,
        s_view, batch, cols, kh,
        depth, 1, sw0, 1,
        stride_p, stride_r, 1,
        ss0, 1,
        BLOCK_SIZE_M=bm, BLOCK_SIZE_N=bn, BLOCK_SIZE_K=bk,
        GROUP_SIZE_M=1, NUM_KSPLIT=ks, SPLITK_BLOCK_SIZE=spk,
        num_warps=4, num_stages=2, waves_per_eu=wpe,
        matrix_instr_nonkdim=16, cache_modifier=".cg",
        PREQUANT=True,
    )

    if frags is not None:
        if _HAS_HIP_MERGER:
            _hip_merger.merge_partials(frags, dest, batch, cols, real_np)
        else:
            _gemm_afp4wfp4_reduce_kernel[red_grid](
                frags, dest, batch, cols,
                batch * cols, cols, 1, cols, 1,
                16, 16, real_np, padded_np,
            )

    return dest if ndims == 2 else dest.view(*X.shape[:-1], cols)
scrolls · 537 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 738857.

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

Best evidence level for this revision: reported

JSON