Skip to content
KernelIndex
Search⌘K

submission 720388

Hamza · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-720388?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.74µs
#84 of 1143
2026-04-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:9817cd905e73fe676072b3279103486928b941c7b3025aa412f4a56a791690a6
license declaredunknown
license concludedunknown
authorsHamza
imported2026-08-15

Techniques

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

fp4print("[hwfp4] Replacing _mxfp4_quant_op with hardware FP4 conversion...", file=_sys.stderr, flush=True)
num-warps = 4num_warps=4, num_stages=2, waves_per_eu=WPE,
split-k_lines = ["cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio"]
stages = 2num_warps=4, num_stages=2, waves_per_eu=WPE,
tile-k = 256BLOCK_K = 256 if K_real <= KSPLIT * 512 or (KSPLIT == 2 and K_real <= KSPLIT * 1024) else 512
tile-m = 8BLOCK_M = 8
tile-n = 128BLOCK_N = 128
vector-width = float4float4 s = *reinterpret_cast<const float4*>(pp + idx4);

Kernel source

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

# submission_v24_hwfp4.py — Hardware FP4 conversion using v_cvt_scalef32_pk_fp4_bf16
# Replaces ~428 ALU quant instructions with ~16 hardware conversion instructions
# Uses tl.inline_asm_elementwise (confirmed available on runner)

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

_KERNEL_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_CSV_PATH = "/tmp/_mxfp4_mm_config.csv"
_CU = 256
_NK_FAMILIES = [
    (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),
]
_M_VALUES = [1, 2, 4, 8, 16, 32, 64, 128, 256]
_lines = ["cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio"]
for _n, _k in _NK_FAMILIES:
    for _m in _M_VALUES:
        _tile_num = ((_m + 31) // 32) * ((_n + 127) // 128)
        _cus_per_tile = _CU / max(_tile_num, 1)
        _split = 0
        while _cus_per_tile >= pow(2, _split + 1) and (pow(2, _split + 1) * 128) < 2 * _k:
            _split += 1
        _split = min(_split, 3)
        _lines.append(f"{_CU},{_m},{_n},{_k},21,{_split},1.0,{_KERNEL_32x128},0,0,0.0")
with open(_CSV_PATH, "w") as _f:
    _f.write("\n".join(_lines))
_os.environ["AITER_CONFIG_GEMM_A4W4"] = _CSV_PATH + ":/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 _sys
import time as _time
import gc as _gc
_sys.setswitchinterval(1.0)

# --- Monkey-patch heuristics ---
try:
    _gemm_a16wfp4_preshuffle_kernel.values['GRID_MN'] = lambda args: 1
    _gemm_a16wfp4_preshuffle_kernel.values['EVEN_K'] = lambda args: True
    print("[patch] GRID_MN → 1, EVEN_K → True", file=_sys.stderr, flush=True)
except Exception as _e:
    print(f"[patch] heuristics failed: {_e}", file=_sys.stderr, flush=True)

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

# --- Replace _mxfp4_quant_op with hardware FP4 conversion ---
print("[hwfp4] Replacing _mxfp4_quant_op with hardware FP4 conversion...", file=_sys.stderr, flush=True)
try:
    _jit_fn = _gemm_a16wfp4_preshuffle_kernel.fn if hasattr(_gemm_a16wfp4_preshuffle_kernel, 'fn') else _gemm_a16wfp4_preshuffle_kernel
    _quant_fn = _jit_fn.__globals__['_mxfp4_quant_op']
    _old_qsrc = _quant_fn._src

    # Complete replacement of _mxfp4_quant_op with hardware FP4 instruction
    _new_qsrc = '''def _mxfp4_quant_op(
    x,
    BLOCK_SIZE_N,
    BLOCK_SIZE_M,
    MXFP4_QUANT_BLOCK_SIZE,
):
    """Hardware-accelerated BF16->MXFP4 using 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)

    # Compute amax per group of 32 (same as original)
    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

    # E8M0 scale computation (v19 integer bit ops)
    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)

    # E8M0 scale bytes for output
    bs_e8m0 = (scale_e8m0_unbiased + 127).to(tl.float32).to(tl.uint8)

    # Hardware scale: DIVISOR (confirmed by probe: scale=0.5 gives fp4(x/0.5)=fp4(2x))
    # Instruction computes: fp4 = round_to_fp4(bf16 / hw_scale)
    # We want: fp4 = round(x / 2^scale_e8m0_unbiased)
    # So hw_scale = 2^scale_e8m0_unbiased, constructed via IEEE 754 bit manipulation
    # biased_exp = scale_unbiased + 127, clamped to [1, 254] (avoid 0 which gives float 0.0)
    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)

    # Convert to BF16 for hardware instruction (x may be float32 from auto-promotion)
    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)  # each [BM, NQ, HALF_BLOCK]
    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)  # [BM, NQ, HALF_BLOCK]

    # Hardware FP4 conversion!
    # hw_scale [BM, NQ, 1] broadcasts to [BM, NQ, HALF_BLOCK] implicitly
    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,
    )

    # Extract byte 0 (the 2 packed FP4 nibbles)
    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_fn, '_unsafe_update_src'):
        _quant_fn._unsafe_update_src(_new_qsrc)
    else:
        _quant_fn._src = _new_qsrc
        if hasattr(_quant_fn, 'src'):
            _quant_fn.src = _new_qsrc
        if hasattr(_quant_fn, 'hash'):
            _quant_fn.hash = None

    # Also modify the KERNEL source to bust its Triton cache key
    _old_ksrc = _jit_fn._src
    _new_ksrc = _old_ksrc.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 _new_ksrc != _old_ksrc:
        _jit_fn._unsafe_update_src(_new_ksrc)
        print("[hwfp4] Applied hardware quant + kernel cache bust", file=_sys.stderr, flush=True)
    else:
        print("[hwfp4] Applied hardware quant, kernel mod FAILED", file=_sys.stderr, flush=True)

    # Verify
    _vq = _quant_fn._src if hasattr(_quant_fn, '_src') else ''
    print(f"[hwfp4] quant has inline_asm: {'inline_asm_elementwise' in _vq}",
          file=_sys.stderr, flush=True)
except Exception as _e:
    import traceback
    print(f"[hwfp4] FAILED: {_e}", file=_sys.stderr, flush=True)
    traceback.print_exc(file=_sys.stderr)

# --- HIP reduce kernel (same as v19) ---
_HIP_REDUCE_SRC = r"""
#include <hip/hip_runtime.h>

__device__ __forceinline__ unsigned short f32_to_bf16(float f) {
    unsigned int u;
    __builtin_memcpy(&u, &f, sizeof(u));
    unsigned int rounding_bias = ((u >> 16) & 1) + 0x7FFFu;
    return (unsigned short)((u + rounding_bias) >> 16);
}

template <int KSPLIT>
__global__ void reduce_k_vec4(const float* __restrict__ pp,
                              unsigned short* __restrict__ out, int MN) {
    int idx4 = (blockIdx.x * blockDim.x + threadIdx.x) * 4;
    if (idx4 + 3 < MN) {
        float4 s = *reinterpret_cast<const float4*>(pp + idx4);
        #pragma unroll
        for (int k = 1; k < KSPLIT; k++) {
            float4 v = *reinterpret_cast<const float4*>(pp + k * MN + idx4);
            s.x += v.x; s.y += v.y; s.z += v.z; s.w += v.w;
        }
        unsigned short r0 = f32_to_bf16(s.x);
        unsigned short r1 = f32_to_bf16(s.y);
        unsigned short r2 = f32_to_bf16(s.z);
        unsigned short r3 = f32_to_bf16(s.w);
        *reinterpret_cast<unsigned long long*>(out + idx4) =
            (unsigned long long)r0 | ((unsigned long long)r1 << 16) |
            ((unsigned long long)r2 << 32) | ((unsigned long long)r3 << 48);
    } else {
        for (int i = idx4; i < MN && i < idx4 + 4; i++) {
            float s = pp[i];
            #pragma unroll
            for (int k = 1; k < KSPLIT; k++) s += pp[k * MN + i];
            out[i] = f32_to_bf16(s);
        }
    }
}

__global__ void reduce_k_gen(const float* __restrict__ pp,
                             unsigned short* __restrict__ out, int MN, int ksplit) {
    int idx = blockIdx.x * blockDim.x + threadIdx.x;
    if (idx < MN) {
        float s = pp[idx];
        for (int k = 1; k < ksplit; k++) s += pp[k * MN + idx];
        out[idx] = f32_to_bf16(s);
    }
}

void reduce_op(torch::Tensor pp, torch::Tensor out, int M, int N, int ksplit) {
    int MN = M * N;
    const float* pp_ptr = pp.data_ptr<float>();
    unsigned short* out_ptr = reinterpret_cast<unsigned short*>(out.data_ptr());
    const int threads_v = 64;
    const int elems_per_block = threads_v * 4;
    const int blocks_v = (MN + elems_per_block - 1) / elems_per_block;
    switch (ksplit) {
        case 2: reduce_k_vec4<2><<<blocks_v, threads_v>>>(pp_ptr, out_ptr, MN); break;
        case 3: reduce_k_vec4<3><<<blocks_v, threads_v>>>(pp_ptr, out_ptr, MN); break;
        case 4: reduce_k_vec4<4><<<blocks_v, threads_v>>>(pp_ptr, out_ptr, MN); break;
        case 7: reduce_k_vec4<7><<<blocks_v, threads_v>>>(pp_ptr, out_ptr, MN); break;
        case 8: reduce_k_vec4<8><<<blocks_v, threads_v>>>(pp_ptr, out_ptr, MN); break;
        default: {
            const int threads = 256;
            const int blocks = (MN + threads - 1) / threads;
            reduce_k_gen<<<blocks, threads>>>(pp_ptr, out_ptr, MN, ksplit);
            break;
        }
    }
}
"""
_HIP_REDUCE_CPP = "void reduce_op(torch::Tensor pp, torch::Tensor out, int M, int N, int ksplit);"

_USE_HIP_REDUCE = False
try:
    from torch.utils.cpp_extension import load_inline as _load_inline
    _hip_reduce_t0 = _time.time()
    _hip_reduce = _load_inline(
        name="mxfp4_reduce_hip",
        cpp_sources=[_HIP_REDUCE_CPP],
        cuda_sources=[_HIP_REDUCE_SRC],
        functions=["reduce_op"],
        verbose=False,
        extra_cuda_cflags=["--offload-arch=gfx950", "-O3"],
    )
    _USE_HIP_REDUCE = True
    print(f"[hip] reduce kernel compiled in {_time.time()-_hip_reduce_t0:.1f}s",
          file=_sys.stderr, flush=True)
except Exception as _e:
    print(f"[hip] reduce kernel FAILED: {_e}", file=_sys.stderr, flush=True)

# --- Helper functions (same as v19) ---
def _get_splitk(K, BLOCK_SIZE_K, NUM_KSPLIT):
    SPLITK_BLOCK_SIZE = (
        triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K
    )
    while NUM_KSPLIT > 1 and BLOCK_SIZE_K > 16:
        if (K % (SPLITK_BLOCK_SIZE // 2) == 0
            and SPLITK_BLOCK_SIZE % BLOCK_SIZE_K == 0
            and K % (BLOCK_SIZE_K // 2) == 0):
            break
        elif K % (SPLITK_BLOCK_SIZE // 2) != 0 and NUM_KSPLIT > 1:
            NUM_KSPLIT = NUM_KSPLIT // 2
        elif SPLITK_BLOCK_SIZE % BLOCK_SIZE_K != 0:
            if NUM_KSPLIT > 1:
                NUM_KSPLIT = NUM_KSPLIT // 2
            elif BLOCK_SIZE_K > 16:
                BLOCK_SIZE_K = BLOCK_SIZE_K // 2
        elif K % (BLOCK_SIZE_K // 2) != 0 and BLOCK_SIZE_K > 16:
            BLOCK_SIZE_K = BLOCK_SIZE_K // 2
        else:
            break
        SPLITK_BLOCK_SIZE = (
            triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K
        )
    return SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT


_CFG_CACHE = {}

def _get_cfg(M, N, K_real):
    key = (M, N, K_real)
    if key in _CFG_CACHE:
        return _CFG_CACHE[key]
    K = K_real // 2
    if M <= 32:
        BLOCK_M = 8
        BLOCK_N = 128
        tiles_128 = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + 127) // 128)
        KSPLIT = 1
        if K_real >= 4096:
            KSPLIT = 7
        elif K_real >= 2048:
            if tiles_128 * 2 >= (_CU * 3) // 4 and tiles_128 * 2 <= _CU:
                KSPLIT = 2
            else:
                KSPLIT = 4
        elif K_real >= 1536:
            if tiles_128 * 2 >= (_CU * 3) // 4 and tiles_128 * 2 <= _CU:
                KSPLIT = 2
            else:
                KSPLIT = 3
        BLOCK_K = 256 if K_real <= KSPLIT * 512 or (KSPLIT == 2 and K_real <= KSPLIT * 1024) else 512
        if tiles_128 * KSPLIT < (_CU * 3) // 4:
            BLOCK_N = 64
        wgs = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + BLOCK_N - 1) // BLOCK_N) * KSPLIT
        cfg = {
            "BLOCK_SIZE_M": BLOCK_M, "BLOCK_SIZE_N": BLOCK_N, "BLOCK_SIZE_K": BLOCK_K,
            "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
            "waves_per_eu": 2 if wgs > _CU else 1, "matrix_instr_nonkdim": 16,
            "cache_modifier": ".cg", "NUM_KSPLIT": KSPLIT,
        }
    else:
        BLOCK_M = 16
        if M <= 128:
            tiles_bm16 = ((M + 15) // 16) * ((N + 127) // 128)
            if tiles_bm16 < (_CU * 3) // 4:
                BLOCK_M = 8
        tiles = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + 127) // 128)
        BLOCK_N = 128
        KSPLIT = 1
        if _CU // 2 <= tiles <= _CU and (K_real >= 7168 or (K_real >= 2048 and BLOCK_M == 8)):
            KSPLIT = 2
        elif tiles < _CU // 2 and K_real > 512:
            if K_real >= 4096:
                if tiles * 2 >= _CU:
                    KSPLIT = 2
                else:
                    KSPLIT = 7
            elif K_real >= 2048:
                KSPLIT = 2
            elif K_real >= 1536:
                KSPLIT = 3
        BLOCK_K = 256 if K_real <= max(KSPLIT * 4096, 2048) else 512
        if tiles * KSPLIT < (_CU * 3) // 4:
            BLOCK_N = 64
        wgs = ((M + BLOCK_M - 1) // BLOCK_M) * ((N + BLOCK_N - 1) // BLOCK_N) * KSPLIT
        cfg = {
            "BLOCK_SIZE_M": BLOCK_M, "BLOCK_SIZE_N": BLOCK_N, "BLOCK_SIZE_K": BLOCK_K,
            "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
            "waves_per_eu": 2 if wgs > _CU else 1, "matrix_instr_nonkdim": 16,
            "cache_modifier": ".cg", "NUM_KSPLIT": KSPLIT,
        }

    if cfg["NUM_KSPLIT"] > 1:
        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"] = 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["SPLITK_BLOCK_SIZE"] = 2 * K

    actual_ksplit = None
    nk_pow2 = None
    if cfg["NUM_KSPLIT"] > 1:
        actual_ksplit = triton.cdiv(K, cfg["SPLITK_BLOCK_SIZE"] // 2)
        nk_pow2 = triton.next_power_of_2(cfg["NUM_KSPLIT"])

    num_m_tiles = triton.cdiv(M, cfg["BLOCK_SIZE_M"])
    num_n_tiles = triton.cdiv(N, cfg["BLOCK_SIZE_N"])
    total_tiles = num_m_tiles * num_n_tiles
    grid_main = (cfg["NUM_KSPLIT"] * total_tiles,)
    grid_reduce = None
    if cfg["NUM_KSPLIT"] > 1:
        grid_reduce = (triton.cdiv(M, 16), triton.cdiv(N, 16))

    result = (cfg, actual_ksplit, nk_pow2, grid_main, grid_reduce,
              K, cfg["BLOCK_SIZE_M"], cfg["BLOCK_SIZE_N"], cfg["BLOCK_SIZE_K"],
              cfg["NUM_KSPLIT"], cfg["SPLITK_BLOCK_SIZE"], cfg["waves_per_eu"])
    _CFG_CACHE[key] = result
    return result

# --- Nuclear pre-warming ---
_WARMUP_T0 = _time.time()
_PREWARMED_CONFIGS = {}
_NO_LSR = {}
_LSR = {}
_REDUCE = set()

for _nw, _kw in _NK_FAMILIES:
    for _mw in _M_VALUES:
        _cw, _aw, _nkw, _, _, _, _, _, _, _, _, _ = _get_cfg(_mw, _nw, _kw)
        _ck = (_cw["BLOCK_SIZE_M"], _cw["BLOCK_SIZE_N"], _cw["BLOCK_SIZE_K"],
               _cw["NUM_KSPLIT"], _cw["SPLITK_BLOCK_SIZE"], _cw["waves_per_eu"])
        if _mw <= 32 and _kw >= 1536:
            _NO_LSR.setdefault(_ck, True)
        else:
            _LSR.setdefault(_ck, True)
        if _aw is not None:
            _REDUCE.add((_aw, _nkw))

for _k in _NO_LSR:
    _LSR.pop(_k, None)

print(f"[pre-warm] {len(_NO_LSR)} no-lsr + {len(_LSR)} lsr GEMM, {len(_REDUCE)} reduce configs",
      file=_sys.stderr, flush=True)

_wA = torch.zeros(32, 8192, dtype=torch.bfloat16, device="cuda")
_wBw = torch.zeros(16, 65536, dtype=torch.uint8, device="cuda")
_wBs = torch.zeros(16, 65536, dtype=torch.uint8, device="cuda")
_wypp = torch.zeros(16, 32, 256, dtype=torch.float32, device="cuda")
_wy = torch.zeros(32, 256, dtype=torch.bfloat16, device="cuda")

def _pw(bm, bn, bk, ks, spk, wpe):
    c = {"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}
    o = _wypp if ks > 1 else _wy
    _gemm_a16wfp4_preshuffle_kernel[(max(ks, 1),)](
        _wA, _wBw, o, _wBs, bm, bn, spk // 2,
        _wA.stride(0), _wA.stride(1), _wBw.stride(0), _wBw.stride(1),
        0 if ks <= 1 else _wypp.stride(0),
        _wy.stride(0) if ks <= 1 else _wypp.stride(1),
        _wy.stride(1) if ks <= 1 else _wypp.stride(2),
        _wBs.stride(0), _wBs.stride(1), PREQUANT=True, **c)

print("[pre-warm] Phase 1: M≤32 K>=1536 (no disable-lsr)...", file=_sys.stderr, flush=True)
for _ck in sorted(_NO_LSR):
    try:
        _pw(*_ck)
        _PREWARMED_CONFIGS[_ck] = "no-lsr"
        print(f"  BM={_ck[0]} BN={_ck[1]} BK={_ck[2]} KS={_ck[3]} SPK={_ck[4]} wpe={_ck[5]} ({_time.time()-_WARMUP_T0:.0f}s)",
              file=_sys.stderr, flush=True)
    except Exception as _e:
        print(f"  {_ck}: FAIL {_e}", file=_sys.stderr, flush=True)

_os.environ["DISABLE_LLVM_OPT"] = "disable-lsr"
print(f"[pre-warm] Phase 2: DISABLE_LLVM_OPT=disable-lsr ({_time.time()-_WARMUP_T0:.0f}s)",
      file=_sys.stderr, flush=True)

_lsr_list = sorted(_LSR)
print(f"[pre-warm] Phase 3: {len(_lsr_list)} remaining GEMM configs...", file=_sys.stderr, flush=True)
for _idx, _ck in enumerate(_lsr_list):
    if _time.time() - _WARMUP_T0 > 200:
        print(f"  timeout — {len(_lsr_list) - _idx} skipped", file=_sys.stderr, flush=True)
        break
    try:
        _pw(*_ck)
        _PREWARMED_CONFIGS[_ck] = "lsr"
        print(f"  BM={_ck[0]} BN={_ck[1]} BK={_ck[2]} KS={_ck[3]} SPK={_ck[4]} wpe={_ck[5]} ({_time.time()-_WARMUP_T0:.0f}s)",
              file=_sys.stderr, flush=True)
    except Exception as _e:
        print(f"  {_ck}: FAIL {_e}", file=_sys.stderr, flush=True)

print(f"[pre-warm] Phase 4: {len(_REDUCE)} reduce configs...", file=_sys.stderr, flush=True)
for _ak, _nk in sorted(_REDUCE):
    if _time.time() - _WARMUP_T0 > 230:
        print("  timeout — remaining skipped", file=_sys.stderr, flush=True)
        break
    try:
        _gemm_afp4wfp4_reduce_kernel[(1, 1)](
            _wypp, _wy, 16, 16,
            _wypp.stride(0), _wypp.stride(1), _wypp.stride(2),
            _wy.stride(0), _wy.stride(1), 16, 16, _ak, _nk)
    except Exception:
        pass

del _wA, _wBw, _wBs, _wypp, _wy, _pw
del _NO_LSR, _LSR, _REDUCE, _lsr_list
torch.cuda.empty_cache()
print(f"[pre-warm] Done: {len(_PREWARMED_CONFIGS)} GEMM configs in {_time.time()-_WARMUP_T0:.0f}s",
      file=_sys.stderr, flush=True)

_gc.disable()

# --- Runtime ---
_PRESHUFFLE_CACHE = {}
_OUT_BUF = {}
_YPP_BUF = {}
_LOGGED = set()

def _get_preshuffle_b(data):
    key = data[3].data_ptr()
    if key not in _PRESHUFFLE_CACHE:
        N = data[3].shape[0]
        K_bytes = data[3].shape[1]
        sm, sn = data[4].shape
        N_groups = N // 32
        B_w = data[3].view(torch.uint8).reshape(N // 16, K_bytes * 16)
        B_s = data[4].view(torch.uint8).reshape(sm // 32, sn * 32)[:N_groups].contiguous()
        _PRESHUFFLE_CACHE[key] = (B_w, B_s, B_w.stride(0), B_s.stride(0))
    return _PRESHUFFLE_CACHE[key]

def custom_kernel(data: input_t) -> output_t:
    A = data[0]
    if not A.is_contiguous():
        A = A.contiguous()
    _ndim = A.ndim
    if _ndim == 2:
        A_2d = A
        M = A.shape[0]
    else:
        A_2d = A.view(-1, A.shape[-1])
        M = A_2d.shape[0]
    N = data[3].shape[0]
    K_bytes = data[3].shape[1]
    K_real = K_bytes * 2

    cfg, actual_ksplit, nk_pow2, grid_main, grid_reduce, K, BM, BN, BK, KS, SPK, WPE = _get_cfg(M, N, K_real)

    _sk = (M, N, K_real)
    if _sk not in _LOGGED:
        _LOGGED.add(_sk)
        print(f"[kernel] M={M} N={N} K={K_real} BM={BM} BN={BN} BK={BK} KS={KS} wpe={WPE} grid={grid_main[0]}",
              file=_sys.stderr, flush=True)

    okey = (M, N)
    if okey not in _OUT_BUF:
        _OUT_BUF[okey] = torch.empty((M, N), dtype=torch.bfloat16, device="cuda")
    y = _OUT_BUF[okey]

    B_w, B_s, stride_bw0, stride_bs0 = _get_preshuffle_b(data)

    if KS > 1:
        ppkey = (nk_pow2, M, N)
        if ppkey not in _YPP_BUF:
            _YPP_BUF[ppkey] = torch.empty((nk_pow2, M, N), dtype=torch.float32, device="cuda")
        y_pp = _YPP_BUF[ppkey]
        stride_ck = M * N
        stride_cm = N
    else:
        y_pp = None
        stride_ck = 0
        stride_cm = N

    _gemm_a16wfp4_preshuffle_kernel[grid_main](
        A_2d, B_w,
        y if y_pp is None else y_pp,
        B_s, M, N, K,
        K_real, 1, stride_bw0, 1,
        stride_ck, stride_cm, 1,
        stride_bs0, 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 y_pp is not None:
        if _USE_HIP_REDUCE:
            _hip_reduce.reduce_op(y_pp, y, M, N, actual_ksplit)
        else:
            _gemm_afp4wfp4_reduce_kernel[grid_reduce](
                y_pp, y, M, N,
                M * N, N, 1, N, 1,
                16, 16, actual_ksplit, nk_pow2,
            )

    if _ndim == 2:
        return y
    return y.view(*A.shape[:-1], N)
scrolls · 563 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 715401.

#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
- # submission_v19_fastquant.py — Optimized E8M0 scale computation (integer bit ops)
- # Replace tl.log2().floor() + tl.exp2() with integer shift/mask in _mxfp4_quant_op
+ # submission_v24_hwfp4.py — Hardware FP4 conversion using v_cvt_scalef32_pk_fp4_bf16
+ # Replaces ~428 ALU quant instructions with ~16 hardware conversion instructions
+ # Uses tl.inline_asm_elementwise (confirmed available on runner)
- # --- Config injection (prevents extra module_gemm_common build ~20s) ---
import os as _os
-
- # Must be set BEFORE torch import for load_inline HIP compilation
_os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
_os.environ.setdefault("CXX", "clang++")
- # Fresh Triton cache to force recompilation with modified quant source
import uuid as _uuid
- _os.environ["TRITON_CACHE_DIR"] = f"/tmp/_triton_fq_{_uuid.uuid4().hex[:8]}"
+ _os.environ["TRITON_CACHE_DIR"] = f"/tmp/_triton_hw_{_uuid.uuid4().hex[:8]}"
_KERNEL_32x128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_CSV_PATH = "/tmp/_mxfp4_mm_config.csv"
⋯ 19 unchanged lines
with open(_CSV_PATH, "w") as _f:
_f.write("\n".join(_lines))
_os.environ["AITER_CONFIG_GEMM_A4W4"] = _CSV_PATH + ":/home/runner/aiter/aiter/configs/a4w4_blockscale_tuned_gemm.csv"
- # --- End config injection ---
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,
)
⋯ 6 unchanged lines
import gc as _gc
_sys.setswitchinterval(1.0)
- # --- Monkey-patch heuristics to constants ---
+ # --- Monkey-patch heuristics ---
try:
_gemm_a16wfp4_preshuffle_kernel.values['GRID_MN'] = lambda args: 1
_gemm_a16wfp4_preshuffle_kernel.values['EVEN_K'] = lambda args: True
print("[patch] GRID_MN → 1, EVEN_K → True", file=_sys.stderr, flush=True)
- except (AttributeError, KeyError, TypeError) as _e:
+ except Exception as _e:
print(f"[patch] heuristics failed: {_e}", file=_sys.stderr, flush=True)
_os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
- # --- Modify _mxfp4_quant_op: replace log2/floor/exp2 with integer bit ops ---
- # The E8M0 scale computation uses tl.log2(amax).floor() - 2, but amax is already
- # a power of 2 (mantissa zeroed by & 0xFF800000). So the exponent can be extracted
- # with integer bit shifts, eliminating expensive v_log_f32 and v_ldexp_f32 SFU instructions.
- print("[fastquant] Modifying _mxfp4_quant_op source...", file=_sys.stderr, flush=True)
+ # --- Replace _mxfp4_quant_op with hardware FP4 conversion ---
+ print("[hwfp4] Replacing _mxfp4_quant_op with hardware FP4 conversion...", file=_sys.stderr, flush=True)
try:
_jit_fn = _gemm_a16wfp4_preshuffle_kernel.fn if hasattr(_gemm_a16wfp4_preshuffle_kernel, 'fn') else _gemm_a16wfp4_preshuffle_kernel
_quant_fn = _jit_fn.__globals__['_mxfp4_quant_op']
- # _mxfp4_quant_op is a JITFunction — do NOT unwrap via .fn
_old_qsrc = _quant_fn._src
- # Replacement 1: log2(amax).floor() → integer bit extraction
- _new_qsrc = _old_qsrc.replace(
- " amax = amax.to(tl.float32, bitcast=True)\n"
- " scale_e8m0_unbiased = tl.log2(amax).floor() - 2",
- " amax_exp = (amax >> 23) & 0xFF\n"
- " scale_e8m0_unbiased = (amax_exp.to(tl.int32) - 129).to(tl.float32)"
- )
+ # Complete replacement of _mxfp4_quant_op with hardware FP4 instruction
+ _new_qsrc = '''def _mxfp4_quant_op(
+ x,
+ BLOCK_SIZE_N,
+ BLOCK_SIZE_M,
+ MXFP4_QUANT_BLOCK_SIZE,
+ ):
+ """Hardware-accelerated BF16->MXFP4 using 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
- # Replacement 2: exp2(-scale) → integer FP32 construction
- _new_qsrc = _new_qsrc.replace(
- " quant_scale = tl.exp2(-scale_e8m0_unbiased)",
- " quant_scale = (((127.0 - scale_e8m0_unbiased).to(tl.int32).to(tl.uint32) << 23)).to(tl.float32, bitcast=True)"
+ x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
+
+ # Compute amax per group of 32 (same as original)
+ 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
+
+ # E8M0 scale computation (v19 integer bit ops)
+ 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)
+
+ # E8M0 scale bytes for output
+ bs_e8m0 = (scale_e8m0_unbiased + 127).to(tl.float32).to(tl.uint8)
+
+ # Hardware scale: DIVISOR (confirmed by probe: scale=0.5 gives fp4(x/0.5)=fp4(2x))
+ # Instruction computes: fp4 = round_to_fp4(bf16 / hw_scale)
+ # We want: fp4 = round(x / 2^scale_e8m0_unbiased)
+ # So hw_scale = 2^scale_e8m0_unbiased, constructed via IEEE 754 bit manipulation
+ # biased_exp = scale_unbiased + 127, clamped to [1, 254] (avoid 0 which gives float 0.0)
+ 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)
+
+ # Convert to BF16 for hardware instruction (x may be float32 from auto-promotion)
+ 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) # each [BM, NQ, HALF_BLOCK]
+ 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) # [BM, NQ, HALF_BLOCK]
+
+ # Hardware FP4 conversion!
+ # hw_scale [BM, NQ, 1] broadcasts to [BM, NQ, HALF_BLOCK] implicitly
+ 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,
)
- if _new_qsrc != _old_qsrc:
- if hasattr(_quant_fn, '_unsafe_update_src'):
- _quant_fn._unsafe_update_src(_new_qsrc)
- else:
- _quant_fn._src = _new_qsrc
- if hasattr(_quant_fn, 'src'):
- _quant_fn.src = _new_qsrc
- if hasattr(_quant_fn, 'hash'):
- _quant_fn.hash = None
- # Also modify the KERNEL source to bust its Triton cache key
- # (quant is a dependency but its hash isn't in the kernel's cache key)
- _old_ksrc = _jit_fn._src
- _new_ksrc = _old_ksrc.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 _new_ksrc != _old_ksrc:
- _jit_fn._unsafe_update_src(_new_ksrc)
- print("[fastquant] Applied quant + kernel source modifications", file=_sys.stderr, flush=True)
- else:
- print("[fastquant] Applied quant mod, kernel mod FAILED", file=_sys.stderr, flush=True)
- # Verify
- _vq = _quant_fn._src if hasattr(_quant_fn, '_src') else ''
- _vk = _jit_fn._src if hasattr(_jit_fn, '_src') else ''
- print(f"[fastquant] quant has amax_exp: {'amax_exp' in _vq}, kernel has acc=: {'acc=accumulator' in _vk}",
- file=_sys.stderr, flush=True)
+ # Extract byte 0 (the 2 packed FP4 nibbles)
+ 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_fn, '_unsafe_update_src'):
+ _quant_fn._unsafe_update_src(_new_qsrc)
else:
- print("[fastquant] WARNING: replacement strings not found — source unchanged", file=_sys.stderr, flush=True)
+ _quant_fn._src = _new_qsrc
+ if hasattr(_quant_fn, 'src'):
+ _quant_fn.src = _new_qsrc
+ if hasattr(_quant_fn, 'hash'):
+ _quant_fn.hash = None
+
+ # Also modify the KERNEL source to bust its Triton cache key
+ _old_ksrc = _jit_fn._src
+ _new_ksrc = _old_ksrc.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 _new_ksrc != _old_ksrc:
+ _jit_fn._unsafe_update_src(_new_ksrc)
+ print("[hwfp4] Applied hardware quant + kernel cache bust", file=_sys.stderr, flush=True)
+ else:
+ print("[hwfp4] Applied hardware quant, kernel mod FAILED", file=_sys.stderr, flush=True)
+
+ # Verify
+ _vq = _quant_fn._src if hasattr(_quant_fn, '_src') else ''
+ print(f"[hwfp4] quant has inline_asm: {'inline_asm_elementwise' in _vq}",
+ file=_sys.stderr, flush=True)
except Exception as _e:
import traceback
- print(f"[fastquant] FAILED: {_e}", file=_sys.stderr, flush=True)
+ print(f"[hwfp4] FAILED: {_e}", file=_sys.stderr, flush=True)
traceback.print_exc(file=_sys.stderr)
- # --- End quant modification ---
-
- # --- HIP reduce kernel ---
+ # --- HIP reduce kernel (same as v19) ---
_HIP_REDUCE_SRC = r"""
#include <hip/hip_runtime.h>
⋯ 64 unchanged lines
}
}
"""
-
_HIP_REDUCE_CPP = "void reduce_op(torch::Tensor pp, torch::Tensor out, int M, int N, int ksplit);"
_USE_HIP_REDUCE = False
⋯ 12 unchanged lines
print(f"[hip] reduce kernel compiled in {_time.time()-_hip_reduce_t0:.1f}s",
file=_sys.stderr, flush=True)
except Exception as _e:
- print(f"[hip] reduce kernel FAILED (using Triton fallback): {_e}",
- file=_sys.stderr, flush=True)
- # --- End HIP reduce kernel ---
+ print(f"[hip] reduce kernel FAILED: {_e}", file=_sys.stderr, flush=True)
-
- # --- Helper functions ---
-
- def _get_splitk(K: int, BLOCK_SIZE_K: int, NUM_KSPLIT: int):
+ # --- Helper functions (same as v19) ---
+ def _get_splitk(K, BLOCK_SIZE_K, NUM_KSPLIT):
SPLITK_BLOCK_SIZE = (
triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K
)
while NUM_KSPLIT > 1 and BLOCK_SIZE_K > 16:
- if (
- K % (SPLITK_BLOCK_SIZE // 2) == 0
+ if (K % (SPLITK_BLOCK_SIZE // 2) == 0
and SPLITK_BLOCK_SIZE % BLOCK_SIZE_K == 0
- and K % (BLOCK_SIZE_K // 2) == 0
- ):
+ and K % (BLOCK_SIZE_K // 2) == 0):
break
elif K % (SPLITK_BLOCK_SIZE // 2) != 0 and NUM_KSPLIT > 1:
NUM_KSPLIT = NUM_KSPLIT // 2
⋯ 12 unchanged lines
return SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT
- _CFG_CACHE: dict = {}
+ _CFG_CACHE = {}
-
- def _get_cfg(M: int, N: int, K_real: int):
+ def _get_cfg(M, N, K_real):
key = (M, N, K_real)
if key in _CFG_CACHE:
return _CFG_CACHE[key]
-
K = K_real // 2
-
if M <= 32:
BLOCK_M = 8
BLOCK_N = 128
⋯ 55 unchanged lines
if cfg["NUM_KSPLIT"] > 1:
SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT = _get_splitk(
- 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
⋯ 27 unchanged lines
_CFG_CACHE[key] = result
return result
-
# --- Nuclear pre-warming ---
_WARMUP_T0 = _time.time()
_PREWARMED_CONFIGS = {}
-
_NO_LSR = {}
_LSR = {}
_REDUCE = set()
⋯ 22 unchanged lines
_wypp = torch.zeros(16, 32, 256, dtype=torch.float32, device="cuda")
_wy = torch.zeros(32, 256, dtype=torch.bfloat16, device="cuda")
-
def _pw(bm, bn, bk, ks, spk, wpe):
c = {"BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": bn, "BLOCK_SIZE_K": bk,
"GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
⋯ 8 unchanged lines
_wy.stride(1) if ks <= 1 else _wypp.stride(2),
_wBs.stride(0), _wBs.stride(1), PREQUANT=True, **c)
-
- # Phase 1: M≤32 K>=1536 without disable-lsr
print("[pre-warm] Phase 1: M≤32 K>=1536 (no disable-lsr)...", file=_sys.stderr, flush=True)
for _ck in sorted(_NO_LSR):
try:
⋯ 4 unchanged lines
except Exception as _e:
print(f" {_ck}: FAIL {_e}", file=_sys.stderr, flush=True)
- # Phase 2: set disable-lsr
_os.environ["DISABLE_LLVM_OPT"] = "disable-lsr"
- print(f"[pre-warm] Phase 2: DISABLE_LLVM_OPT=disable-lsr set ({_time.time()-_WARMUP_T0:.0f}s)",
+ print(f"[pre-warm] Phase 2: DISABLE_LLVM_OPT=disable-lsr ({_time.time()-_WARMUP_T0:.0f}s)",
file=_sys.stderr, flush=True)
- # Phase 3: remaining GEMM configs with disable-lsr
_lsr_list = sorted(_LSR)
- print(f"[pre-warm] Phase 3: {len(_lsr_list)} remaining GEMM configs (disable-lsr)...",
- file=_sys.stderr, flush=True)
+ print(f"[pre-warm] Phase 3: {len(_lsr_list)} remaining GEMM configs...", file=_sys.stderr, flush=True)
for _idx, _ck in enumerate(_lsr_list):
if _time.time() - _WARMUP_T0 > 200:
- print(f" timeout safety — {len(_lsr_list) - _idx} configs skipped",
- file=_sys.stderr, flush=True)
+ print(f" timeout — {len(_lsr_list) - _idx} skipped", file=_sys.stderr, flush=True)
break
try:
_pw(*_ck)
⋯ 3 unchanged lines
except Exception as _e:
print(f" {_ck}: FAIL {_e}", file=_sys.stderr, flush=True)
- # Phase 4: reduce kernel configs
print(f"[pre-warm] Phase 4: {len(_REDUCE)} reduce configs...", file=_sys.stderr, flush=True)
for _ak, _nk in sorted(_REDUCE):
if _time.time() - _WARMUP_T0 > 230:
- print(" timeout safety — remaining reduce configs skipped", file=_sys.stderr, flush=True)
+ print(" timeout — remaining skipped", file=_sys.stderr, flush=True)
break
try:
_gemm_afp4wfp4_reduce_kernel[(1, 1)](
_wypp, _wy, 16, 16,
_wypp.stride(0), _wypp.stride(1), _wypp.stride(2),
_wy.stride(0), _wy.stride(1), 16, 16, _ak, _nk)
- print(f" ksplit={_ak} nk_pow2={_nk} ({_time.time()-_WARMUP_T0:.0f}s)",
- file=_sys.stderr, flush=True)
- except Exception as _e:
- print(f" ksplit={_ak} nk={_nk}: FAIL {_e}", file=_sys.stderr, flush=True)
+ except Exception:
+ pass
del _wA, _wBw, _wBs, _wypp, _wy, _pw
del _NO_LSR, _LSR, _REDUCE, _lsr_list
⋯ 2 unchanged lines
file=_sys.stderr, flush=True)
_gc.disable()
- # --- End pre-warming ---
+ # --- Runtime ---
+ _PRESHUFFLE_CACHE = {}
+ _OUT_BUF = {}
+ _YPP_BUF = {}
+ _LOGGED = set()
- _PRESHUFFLE_CACHE: dict = {}
- _OUT_BUF: dict = {}
- _YPP_BUF: dict = {}
- _LOGGED: set = set()
-
-
def _get_preshuffle_b(data):
key = data[3].data_ptr()
if key not in _PRESHUFFLE_CACHE:
⋯ 6 unchanged lines
_PRESHUFFLE_CACHE[key] = (B_w, B_s, B_w.stride(0), B_s.stride(0))
return _PRESHUFFLE_CACHE[key]
-
def custom_kernel(data: input_t) -> output_t:
A = data[0]
if not A.is_contiguous():
A = A.contiguous()
-
_ndim = A.ndim
if _ndim == 2:
A_2d = A
⋯ 10 unchanged lines
_sk = (M, N, K_real)
if _sk not in _LOGGED:
_LOGGED.add(_sk)
- print(f"[kernel] M={M} N={N} K={K_real} BM={BM} BN={BN} "
- f"BK={BK} KS={KS} wpe={WPE} "
- f"grid={grid_main[0]}", file=_sys.stderr, flush=True)
+ print(f"[kernel] M={M} N={N} K={K_real} BM={BM} BN={BN} BK={BK} KS={KS} wpe={WPE} grid={grid_main[0]}",
+ file=_sys.stderr, flush=True)
okey = (M, N)
if okey not in _OUT_BUF:
⋯ 5 unchanged lines
if KS > 1:
ppkey = (nk_pow2, M, N)
if ppkey not in _YPP_BUF:
- _YPP_BUF[ppkey] = torch.empty(
- (nk_pow2, M, N), dtype=torch.float32, device="cuda"
- )
+ _YPP_BUF[ppkey] = torch.empty((nk_pow2, M, N), dtype=torch.float32, device="cuda")
y_pp = _YPP_BUF[ppkey]
stride_ck = M * N
stride_cm = N
⋯ 5 unchanged lines
_gemm_a16wfp4_preshuffle_kernel[grid_main](
A_2d, B_w,
y if y_pp is None else y_pp,
- B_s,
- M, N, K,
- K_real, 1,
- stride_bw0, 1,
+ B_s, M, N, K,
+ K_real, 1, stride_bw0, 1,
stride_ck, stride_cm, 1,
stride_bs0, 1,
BLOCK_SIZE_M=BM, BLOCK_SIZE_N=BN, BLOCK_SIZE_K=BK,
⋯ 9 unchanged lines
else:
_gemm_afp4wfp4_reduce_kernel[grid_reduce](
y_pp, y, M, N,
- M * N, N, 1,
- N, 1,
- 16, 16,
- actual_ksplit, nk_pow2,
+ M * N, N, 1, N, 1,
+ 16, 16, actual_ksplit, nk_pow2,
)
if _ndim == 2:
scrolls · 426 diff lines total

Best evidence level for this revision: reported

JSON