Skip to content
KernelIndex
Search⌘K

submission 550508

div22 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

solution_new_25t_1.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-550508?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
13.1µs
#406 of 1143
2026-03-14

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a9af7d6037fb569e7ed901880f443b1832e1295be4d51f8e010270d3815ee5cf
license declaredunknown
license concludedunknown
authorsdiv22
imported2026-08-15

Techniques

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

fp4MXFP4 GEMM v25t_1 — v25t + two-phase exhaustive config search on gfx950 (MI355X).
split-kand (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)

Kernel source

solution_new_25t_1.py641 lines
"""
MXFP4 GEMM v25t_1 — v25t + two-phase exhaustive config search on gfx950 (MI355X).

Phase 1: sweep tile params (BM, BN, BK, NUM_KSPLIT) with default (warps, stages, etc.)
Phase 2: sweep (warps, stages, wpe, gsm, cm) around best tile config.

Edit SEARCH_SHAPES below to control which benchmark shapes to search per submission.
"""

# ============================================================================
# EDIT THIS: which benchmark shapes to search this run.
# Comment/uncomment to split across multiple submissions (~5-12 min each).
# ============================================================================
SEARCH_SHAPES = {
    #(4,   2880,  512),   # ~3-4 min
    (16,  2112, 7168),   # ~8-12 min
    # (32,  4096,  512),   # ~5-7 min
    # (32,  2880,  512),   # ~5-7 min
    # (64,  7168, 2048),   # ~8-12 min
    # (256, 3072, 1536),   # ~9-13 min
}
import os
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"

import itertools
from typing import Tuple
import torch
from torch.utils.cpp_extension import load_inline
import triton
import triton.language as tl
import uuid


# ---------------------------------------------------------------------------
# C++ quant kernel (from v25d)
# ---------------------------------------------------------------------------

HIP_KERNEL = r"""
#include <hip/hip_runtime.h>
#include <stdint.h>

using bf16x2 = __bf16 __attribute__((ext_vector_type(2)));

__device__ __forceinline__ uint8_t hw_bf16x2_to_fp4x2(uint32_t bf16_pair, float scale) {
    uint32_t result;
    asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
                 : "=v"(result) : "v"(bf16_pair), "v"(scale));
    return static_cast<uint8_t>(result & 0xFFu);
}

__global__ void __launch_bounds__(128, 4)
mxfp4_quant(
    const __bf16* __restrict__ A_bf16,
    uint8_t*      __restrict__ A_fp4,
    uint8_t*      __restrict__ A_scale,
    int M, int K)
{
    const int KS    = K / 32;
    const int K2    = K / 2;
    const int group = blockIdx.x * 128 + threadIdx.x;
    const int row   = group / KS;
    const int kg    = group % KS;
    if (row >= M) return;

    const auto* src = A_bf16 + (long)row * K + kg * 32;
    float absMax = 1e-10f;
    #pragma unroll
    for (int i = 0; i < 32; ++i) {
        float v = __builtin_elementwise_abs(static_cast<float>(src[i]));
        absMax = (v > absMax) ? v : absMax;
    }

    uint32_t u32     = __builtin_bit_cast(uint32_t, absMax);
    const uint32_t amax_exp = ((u32 + 0x200000u) >> 23) & 0xFFu;
    const uint32_t inv_exp  = (amax_exp >= 2u) ? (amax_exp - 2u) : 0u;
    A_scale[(long)row * KS + kg] = static_cast<uint8_t>(inv_exp);

    const float hw_scale = __builtin_bit_cast(float, static_cast<uint32_t>(inv_exp) << 23);
    const uint32_t* src_u32 = reinterpret_cast<const uint32_t*>(src);
    auto* dst = reinterpret_cast<uint8_t*>(A_fp4 + (long)row * K2 + kg * 16);
    #pragma unroll
    for (int i = 0; i < 16; ++i) {
        dst[i] = hw_bf16x2_to_fp4x2(src_u32[i], hw_scale);
    }
}

extern "C" void launch_quant(
    const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale, int M, int K)
{
    const int KS       = K / 32;
    const int n_groups = M * KS;
    const dim3 block{128};
    const dim3 grid{static_cast<uint32_t>((n_groups + 127) / 128)};
    mxfp4_quant<<<grid, block>>>(A_bf16, A_fp4, A_scale, M, K);
}
"""

CPP = r"""
#include <torch/extension.h>
#include <c10/core/DeviceGuard.h>

extern "C" void launch_quant(const __bf16*, uint8_t*, uint8_t*, int, int);

struct QuantWorkspace {
    at::Tensor A_fp4;
    at::Tensor A_scale;
    int64_t last_M = -1, last_K = -1;
    void ensure(int M, int K, const at::TensorOptions& opts) {
        if (M == last_M && K == last_K) return;
        A_fp4   = at::empty({(int64_t)M, (int64_t)(K / 2)}, opts.dtype(at::kByte));
        A_scale = at::empty({(int64_t)M, (int64_t)(K / 32)}, opts.dtype(at::kByte));
        last_M = M; last_K = K;
    }
};

static QuantWorkspace g_qws;

std::vector<at::Tensor> do_quant(const at::Tensor& A) {
    auto guard = at::DeviceGuard(A.device());
    at::Tensor A_bf16 = (A.scalar_type() == at::kBFloat16 && A.is_contiguous())
                        ? A : A.to(A.device(), at::kBFloat16, false, false,
                                   at::MemoryFormat::Contiguous);
    const int M = A_bf16.size(0);
    const int K = A_bf16.size(1);
    g_qws.ensure(M, K, A_bf16.options());
    launch_quant(
        reinterpret_cast<const __bf16*>(A_bf16.data_ptr<at::BFloat16>()),
        g_qws.A_fp4.data_ptr<uint8_t>(),
        g_qws.A_scale.data_ptr<uint8_t>(),
        M, K);
    return {g_qws.A_fp4, g_qws.A_scale};
}
"""

_ext = load_inline(
    name=f"g_{uuid.uuid4().hex[:8]}",
    cpp_sources=[CPP],
    cuda_sources=[HIP_KERNEL],
    functions=["do_quant"],
    with_cuda=True,
    extra_cflags=["-O3", "-std=c++20"],
    extra_cuda_cflags=[
        "-O3", "--offload-arch=gfx950", "-ffast-math", "-munsafe-fp-atomics",
        "-std=c++20", "-mllvm", "-amdgpu-early-inline-all=true",
        "-mllvm", "-amdgpu-function-calls=false", "-mwavefrontsize64",
        "-mcumode", "-fgpu-flush-denormals-to-zero",
    ],
    extra_ldflags=["-lamdhip64"],
)


# ---------------------------------------------------------------------------
# Triton helpers
# ---------------------------------------------------------------------------

@triton.jit
def remap_xcd(pid, GRID_MN, NUM_XCDS: tl.constexpr = 8):
    pids_per_xcd = (GRID_MN + NUM_XCDS - 1) // NUM_XCDS
    tall_xcds = GRID_MN % NUM_XCDS
    tall_xcds = NUM_XCDS if tall_xcds == 0 else tall_xcds
    xcd = pid % NUM_XCDS
    local_pid = pid // NUM_XCDS
    if xcd < tall_xcds:
        pid = xcd * pids_per_xcd + local_pid
    else:
        pid = tall_xcds * pids_per_xcd + (xcd - tall_xcds) * (pids_per_xcd - 1) + local_pid
    return pid


@triton.jit
def pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M: tl.constexpr = 1):
    if GROUP_SIZE_M == 1:
        pid_m = pid // num_pid_n
        pid_n = pid % num_pid_n
    else:
        num_pid_in_group = GROUP_SIZE_M * num_pid_n
        group_id = pid // num_pid_in_group
        first_pid_m = group_id * GROUP_SIZE_M
        group_size_m = tl.minimum(num_pid_m - first_pid_m, GROUP_SIZE_M)
        pid_m = first_pid_m + (pid % group_size_m)
        pid_n = (pid % num_pid_in_group) // group_size_m
    return pid_m, pid_n


# ---------------------------------------------------------------------------
# Triton GEMM kernel (identical to v25t)
# ---------------------------------------------------------------------------

@triton.heuristics({
    "EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)
    and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
    and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),
})
@triton.jit
def _mxfp4_gemm_kernel(
    a_ptr, b_ptr, c_ptr, a_scales_ptr, b_scales_ptr,
    M, N, K,
    stride_am, stride_ak, stride_bn, stride_bk,
    stride_ck, stride_cm, stride_cn,
    stride_asm, stride_ask, stride_bsn, stride_bsk,
    BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
    BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr,
    NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,
    EVEN_K: tl.constexpr, num_warps: tl.constexpr,
    num_stages: tl.constexpr, waves_per_eu: tl.constexpr,
    matrix_instr_nonkdim: tl.constexpr, cache_modifier: tl.constexpr,
):
    tl.assume(stride_am > 0); tl.assume(stride_ak > 0)
    tl.assume(stride_bk > 0); tl.assume(stride_bn > 0)
    tl.assume(stride_cm > 0); tl.assume(stride_cn > 0)
    tl.assume(stride_asm > 0); tl.assume(stride_ask > 0)
    tl.assume(stride_bsk > 0); tl.assume(stride_bsn > 0)

    GRID_MN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)
    SCALE_GROUP_SIZE: tl.constexpr = 32

    pid_unified = tl.program_id(axis=0)
    pid_unified = remap_xcd(pid_unified, GRID_MN * NUM_KSPLIT, NUM_XCDS=8)
    pid_k = pid_unified % NUM_KSPLIT
    pid = pid_unified // NUM_KSPLIT
    num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
    num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)

    if NUM_KSPLIT == 1:
        pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
    else:
        pid_m = pid // num_pid_n
        pid_n = pid % num_pid_n

    tl.assume(pid_m >= 0); tl.assume(pid_n >= 0)

    if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
        num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)

        offs_k = tl.arange(0, BLOCK_SIZE_K // 2)
        offs_k_split = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_k
        offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
        a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_k_split[None, :] * stride_ak

        offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
        offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
        offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
        b_ptrs = b_ptr + offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk

        offs_ks_a = pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) + tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
        a_scale_ptrs = a_scales_ptr + offs_am[:, None] * stride_asm + offs_ks_a[None, :] * stride_ask

        offs_asn = (pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)) % N
        offs_ks_b = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)
        b_scale_ptrs = b_scales_ptr + offs_asn[:, None] * stride_bsn + offs_ks_b[None, :] * stride_bsk

        accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)

        for k in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
            a_scales = tl.load(a_scale_ptrs)
            b_scales = (
                tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
                .reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8, 4, 16, 2, 2, 1)
                .permute(0, 5, 3, 1, 4, 2, 6)
                .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
            )
            if EVEN_K:
                a = tl.load(a_ptrs)
                b = tl.load(b_ptrs, cache_modifier=cache_modifier)
            else:
                a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k * (BLOCK_SIZE_K // 2), other=0)
                b = tl.load(b_ptrs, mask=offs_k_shuffle_arr[None, :] < (K - k * (BLOCK_SIZE_K // 2)) * 16, other=0, cache_modifier=cache_modifier)

            b = b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16).permute(0, 1, 4, 2, 3, 5).reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2).trans(1, 0)
            accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)

            a_ptrs += (BLOCK_SIZE_K // 2) * stride_ak
            b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
            a_scale_ptrs += (BLOCK_SIZE_K // SCALE_GROUP_SIZE) * stride_ask
            b_scale_ptrs += BLOCK_SIZE_K * stride_bsk

        c = accumulator.to(c_ptr.type.element_ty)
        offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
        offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
        c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :] + pid_k * stride_ck
        c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
        tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")


@triton.jit
def _mxfp4_reduce_kernel(
    c_in_ptr, c_out_ptr, M, N,
    stride_c_in_k, stride_c_in_m, stride_c_in_n,
    stride_c_out_m, stride_c_out_n,
    BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
    ACTUAL_KSPLIT: tl.constexpr, MAX_KSPLIT: tl.constexpr,
):
    pid_m = tl.program_id(axis=0)
    pid_n = tl.program_id(axis=1)
    offs_m = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
    offs_n = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
    offs_k = tl.arange(0, MAX_KSPLIT)
    c_in_ptrs = c_in_ptr + offs_k[:, None, None] * stride_c_in_k + offs_m[None, :, None] * stride_c_in_m + offs_n[None, None, :] * stride_c_in_n
    if ACTUAL_KSPLIT == MAX_KSPLIT:
        c = tl.load(c_in_ptrs)
    else:
        c = tl.load(c_in_ptrs, mask=offs_k[:, None, None] < ACTUAL_KSPLIT)
    c = tl.sum(c, axis=0).to(c_out_ptr.type.element_ty)
    c_out_ptrs = c_out_ptr + offs_m[:, None] * stride_c_out_m + offs_n[None, :] * stride_c_out_n
    tl.store(c_out_ptrs, c)


# ---------------------------------------------------------------------------
# AITER default configs (fallback / seeds)
# ---------------------------------------------------------------------------

def _cfg(bm, bn, bk, gsm, nw, ns, wpe, cm, nks):
    return {
        "BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": bn, "BLOCK_SIZE_K": bk,
        "GROUP_SIZE_M": gsm, "num_warps": nw, "num_stages": ns,
        "waves_per_eu": wpe, "matrix_instr_nonkdim": 16,
        "cache_modifier": cm, "NUM_KSPLIT": nks,
    }

_DEFAULTS = {
    (4,  2880, 512):   _cfg(8,  64, 512, 1, 2, 1, 1, None, 1),
    (16, 2112, 7168):  _cfg(16, 32, 512, 1, 4, 1, 4, ".cg", 14),
    (32, 4096, 512):   _cfg(32, 64, 512, 1, 4, 1, 1, None, 1),
    (32, 2880, 512):   _cfg(32, 64, 512, 1, 4, 1, 1, None, 1),
    (64, 7168, 2048):  _cfg(32, 64, 512, 4, 2, 2, 1, None, 1),
    (256,3072, 1536):  _cfg(128,32, 512, 1, 4, 2, 1, None, 1),
    # test shapes
    (8,  2112, 7168):  _cfg(8,  32, 512, 1, 2, 2, 1, ".cg", 7),
    (16, 3072, 1536):  _cfg(8,  32, 512, 1, 4, 2, 1, None, 1),
    (64, 3072, 1536):  _cfg(64, 32, 512, 1, 2, 2, 1, ".cg", 1),
    (256,2880, 512):   _cfg(32, 64, 512, 1, 2, 2, 1, None, 1),
}


# ---------------------------------------------------------------------------
# splitK helper (from AITER)
# ---------------------------------------------------------------------------

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
    NUM_KSPLIT = triton.cdiv(K, SPLITK_BLOCK_SIZE // 2)
    return SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT


# ---------------------------------------------------------------------------
# Kernel runner: launches GEMM (+ reduce) with a given config
# ---------------------------------------------------------------------------

def _run_gemm(config, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed, y, y_pp):
    cfg = dict(config)
    BK = cfg["BLOCK_SIZE_K"]
    NUM_KSPLIT = cfg["NUM_KSPLIT"]

    if BK >= 2 * K_packed:
        BK = triton.next_power_of_2(2 * K_packed)
        cfg["BLOCK_SIZE_K"] = BK
        cfg["NUM_KSPLIT"] = 1
        NUM_KSPLIT = 1

    cfg["BLOCK_SIZE_K"] = max(cfg["BLOCK_SIZE_K"], 256)
    BK = cfg["BLOCK_SIZE_K"]

    if NUM_KSPLIT > 1:
        SPLITK_BS, BK, NUM_KSPLIT = get_splitk(K_packed, BK, NUM_KSPLIT)
        cfg["BLOCK_SIZE_K"] = BK
        cfg["NUM_KSPLIT"] = NUM_KSPLIT
        cfg["SPLITK_BLOCK_SIZE"] = SPLITK_BS
    else:
        cfg["SPLITK_BLOCK_SIZE"] = 2 * K_packed

    c_out = y_pp if NUM_KSPLIT > 1 else y

    KS = K_elem // 32
    scaleN = ((KS + 7) // 8) * 8

    grid = lambda META: (
        META["NUM_KSPLIT"]
        * triton.cdiv(M, META["BLOCK_SIZE_M"])
        * triton.cdiv(N, META["BLOCK_SIZE_N"]),
    )

    _mxfp4_gemm_kernel[grid](
        A_fp4, w, c_out, A_scale, b_scales,
        M, N, K_packed,
        A_fp4.stride(0), A_fp4.stride(1),
        w.stride(0), w.stride(1),
        0 if NUM_KSPLIT == 1 else y_pp.stride(0),
        c_out.stride(-2), c_out.stride(-1),
        A_scale.stride(0), A_scale.stride(1),
        32 * scaleN, 1,
        **cfg,
    )

    if NUM_KSPLIT > 1:
        ACTUAL_KSPLIT = triton.cdiv(K_packed, cfg["SPLITK_BLOCK_SIZE"] // 2)
        grid_r = (triton.cdiv(M, 16), triton.cdiv(N, 64))
        _mxfp4_reduce_kernel[grid_r](
            y_pp, y, M, N,
            y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
            y.stride(0), y.stride(1),
            16, 64, ACTUAL_KSPLIT, triton.next_power_of_2(NUM_KSPLIT),
        )

    return y


# ---------------------------------------------------------------------------
# Two-phase exhaustive config search
# ---------------------------------------------------------------------------

_CONFIG_CACHE = {}


def _validate_config(config, K_packed):
    BN = config["BLOCK_SIZE_N"]
    BK = config["BLOCK_SIZE_K"]
    if BN < 32:
        return False
    if BK < 64:
        return False
    if K_packed % (BK // 2) != 0:
        return False
    return True


def _time_config(config, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed,
                 n_warmup=2, n_iter=8):
    NUM_KSPLIT = config.get("NUM_KSPLIT", 1)
    y = torch.empty((M, N), dtype=torch.bfloat16, device=A_fp4.device)
    y_pp = None
    if NUM_KSPLIT > 1:
        y_pp = torch.empty((NUM_KSPLIT, M, N), dtype=torch.float32, device=A_fp4.device)

    for _ in range(n_warmup):
        _run_gemm(config, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed, y, y_pp)

    torch.cuda.synchronize()
    start_evt = torch.cuda.Event(enable_timing=True)
    end_evt = torch.cuda.Event(enable_timing=True)
    start_evt.record()
    for _ in range(n_iter):
        _run_gemm(config, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed, y, y_pp)
    end_evt.record()
    end_evt.synchronize()

    return start_evt.elapsed_time(end_evt) / n_iter * 1000  # ms -> us


def _get_param_choices(M, N, K_elem):
    K_packed = K_elem // 2

    bm_choices = [b for b in [8, 16, 32, 64, 128, 256] if b <= max(M * 2, 8)]
    bn_choices = [b for b in [32, 64, 128, 256] if b <= N]
    bk_choices = [b for b in [256, 512, 1024] if K_packed % (b // 2) == 0]
    if not bk_choices:
        bk_choices = [256]

    max_splits = K_packed // (min(bk_choices) // 2)
    sk_choices = [1]
    for s in [2, 3, 4, 7, 14]:
        if s <= max_splits and K_packed % s == 0:
            sk_choices.append(s)

    return bm_choices, bn_choices, bk_choices, sk_choices


def _search_config(M, N, K_elem, A_fp4, A_scale, w, b_scales):
    K_packed = K_elem // 2
    bm_choices, bn_choices, bk_choices, sk_choices = _get_param_choices(M, N, K_elem)

    # Phase 1: sweep tile params (BM, BN, BK, NUM_KSPLIT) with sensible defaults
    tile_combos = list(itertools.product(bm_choices, bn_choices, bk_choices, sk_choices))
    p1_total = len(tile_combos)
    print(f"  Phase 1: {p1_total} tile combos (BM x BN x BK x SK = "
          f"{len(bm_choices)}x{len(bn_choices)}x{len(bk_choices)}x{len(sk_choices)})", flush=True)

    best_time = float("inf")
    best_tile = None
    best_config = _DEFAULTS.get((M, N, K_elem), _cfg(16, 64, 256, 4, 2, 2, 1, None, 1))

    for i, (bm, bn, bk, sk) in enumerate(tile_combos):
        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": 1, "matrix_instr_nonkdim": 16,
            "cache_modifier": None, "NUM_KSPLIT": sk,
        }
        if not _validate_config(cfg, K_packed):
            continue
        try:
            t = _time_config(cfg, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed)
        except Exception:
            continue
        if t < best_time:
            best_time = t
            best_tile = (bm, bn, bk, sk)
            best_config = cfg
            print(f"  [P1 {i+1}/{p1_total}] best={t:.1f}us BM={bm} BN={bn} BK={bk} SK={sk}", flush=True)

    if best_tile is None:
        print("  Phase 1: no valid tile found, using default", flush=True)
        return best_config

    bm, bn, bk, sk = best_tile
    print(f"  Phase 1 winner: BM={bm} BN={bn} BK={bk} SK={sk} = {best_time:.1f}us", flush=True)

    # Phase 2: sweep (num_warps, num_stages, waves_per_eu, GROUP_SIZE_M, cache_modifier)
    tune_combos = list(itertools.product(
        [2, 4],           # num_warps
        [1, 2],           # num_stages
        [1, 2, 4],        # waves_per_eu
        [1, 4, 8],        # GROUP_SIZE_M
        [None, ".cg"],    # cache_modifier
    ))
    p2_total = len(tune_combos)
    print(f"  Phase 2: {p2_total} tune combos around best tile", flush=True)

    for i, (nw, ns, wpe, gsm, cm) in enumerate(tune_combos):
        cfg = {
            "BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": bn, "BLOCK_SIZE_K": bk,
            "GROUP_SIZE_M": gsm, "num_warps": nw, "num_stages": ns,
            "waves_per_eu": wpe, "matrix_instr_nonkdim": 16,
            "cache_modifier": cm, "NUM_KSPLIT": sk,
        }
        try:
            t = _time_config(cfg, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed)
        except Exception:
            continue
        if t < best_time:
            best_time = t
            best_config = cfg
            print(f"  [P2 {i+1}/{p2_total}] best={t:.1f}us nw={nw} ns={ns} wpe={wpe} gsm={gsm} cm={cm}", flush=True)

    print(f"  Phase 2 winner: {best_time:.1f}us", flush=True)
    return best_config


# ---------------------------------------------------------------------------
# Workspace caching
# ---------------------------------------------------------------------------

class _Workspace:
    __slots__ = ["y", "y_pp", "_key"]
    def __init__(self):
        self.y = None; self.y_pp = None; self._key = None
    def ensure(self, M, N, num_ksplit, device):
        key = (M, N, num_ksplit)
        if key == self._key: return
        self._key = key
        self.y = torch.empty((M, N), dtype=torch.bfloat16, device=device)
        if num_ksplit > 1:
            self.y_pp = torch.empty((num_ksplit, M, N), dtype=torch.float32, device=device)
        else:
            self.y_pp = None

_ws = _Workspace()
_b_cache = {"bsh_dp": 0, "bssh_dp": 0, "w": None, "b_scales": None}


# ---------------------------------------------------------------------------
# Main entry point
# ---------------------------------------------------------------------------

def custom_kernel(data: Tuple[torch.Tensor, ...]) -> torch.Tensor:
    """MXFP4 GEMM v25t_1: two-phase exhaustive config search."""
    A, _, B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous().cuda()

    M, K_elem = A.shape
    N = B_q.shape[0]
    K_packed = K_elem // 2

    # Quant
    quant_result = _ext.do_quant(A)
    A_fp4 = quant_result[0]
    A_scale = quant_result[1]

    # Cache B views
    bsh_dp = B_shuffle.data_ptr()
    bssh_dp = B_scale_sh.data_ptr()
    if bsh_dp != _b_cache["bsh_dp"] or bssh_dp != _b_cache["bssh_dp"]:
        w_raw = B_shuffle.view(torch.uint8) if B_shuffle.dtype != torch.uint8 else B_shuffle
        _b_cache["w"] = w_raw.reshape(N // 16, K_packed * 16).contiguous()
        bs_raw = B_scale_sh.view(torch.uint8) if B_scale_sh.dtype != torch.uint8 else B_scale_sh
        _b_cache["b_scales"] = bs_raw.contiguous()
        _b_cache["bsh_dp"] = bsh_dp
        _b_cache["bssh_dp"] = bssh_dp

    w = _b_cache["w"]
    b_scales = _b_cache["b_scales"]

    # --- Search only for shapes listed in SEARCH_SHAPES; defaults for rest ---
    shape_key = (M, N, K_elem)
    if shape_key not in _CONFIG_CACHE:
        if shape_key in SEARCH_SHAPES:
            print(f"[v25t_1] exhaustive search for M={M} N={N} K={K_elem} ...", flush=True)
            try:
                best = _search_config(M, N, K_elem, A_fp4, A_scale, w, b_scales)
                _CONFIG_CACHE[shape_key] = best
                cm = best.get("cache_modifier", None)
                print(f"[v25t_1] BEST M={M} N={N} K={K_elem}: "
                      f"BM={best['BLOCK_SIZE_M']} BN={best['BLOCK_SIZE_N']} "
                      f"BK={best['BLOCK_SIZE_K']} warps={best['num_warps']} "
                      f"stages={best['num_stages']} wpe={best['waves_per_eu']} "
                      f"GSM={best['GROUP_SIZE_M']} splitK={best['NUM_KSPLIT']} "
                      f"cache={cm}", flush=True)
            except Exception as e:
                print(f"[v25t_1] search failed: {e}, using default", flush=True)
                _CONFIG_CACHE[shape_key] = _DEFAULTS.get(
                    shape_key, _cfg(16, 64, 256, 4, 2, 2, 1, None, 1)
                )
        else:
            # Test shape — use defaults, no search
            _CONFIG_CACHE[shape_key] = _DEFAULTS.get(
                shape_key, _cfg(16, 64, 256, 4, 2, 2, 1, None, 1)
            )
            print(f"[v25t_1] test shape M={M} N={N} K={K_elem}, using default", flush=True)

    config = _CONFIG_CACHE[shape_key]
    NUM_KSPLIT = config.get("NUM_KSPLIT", 1)

    _ws.ensure(M, N, NUM_KSPLIT, A.device)

    return _run_gemm(config, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed,
                     _ws.y, _ws.y_pp)
scrolls · 641 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 546301.

"""
- MXFP4 GEMM v25d — gfx950 (MI355X) optimized.
+ MXFP4 GEMM v25t_1 — v25t + two-phase exhaustive config search on gfx950 (MI355X).
- Changes from v25b (13.459μs):
- 1. Full (N,K) template specialization + precomputed views (same as v25b)
- 2. Aggressive compiler flags:
- - -amdgpu-loop-prefetch: software prefetch for K-loop loads
- - -enable-unroll-and-jam: fuse nested loop unrolling
- - -ffinite-math-only: assume no NaN/Inf (beyond -ffast-math)
- - -amdgpu-set-wave-priority: dynamic wave priority
- - Increased unroll thresholds
+ Phase 1: sweep tile params (BM, BN, BK, NUM_KSPLIT) with default (warps, stages, etc.)
+ Phase 2: sweep (warps, stages, wpe, gsm, cm) around best tile config.
+
+ Edit SEARCH_SHAPES below to control which benchmark shapes to search per submission.
"""
+
+ # ============================================================================
+ # EDIT THIS: which benchmark shapes to search this run.
+ # Comment/uncomment to split across multiple submissions (~5-12 min each).
+ # ============================================================================
+ SEARCH_SHAPES = {
+ #(4, 2880, 512), # ~3-4 min
+ (16, 2112, 7168), # ~8-12 min
+ # (32, 4096, 512), # ~5-7 min
+ # (32, 2880, 512), # ~5-7 min
+ # (64, 7168, 2048), # ~8-12 min
+ # (256, 3072, 1536), # ~9-13 min
+ }
import os
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
+ import itertools
from typing import Tuple
import torch
from torch.utils.cpp_extension import load_inline
+ import triton
+ import triton.language as tl
import uuid
+ # ---------------------------------------------------------------------------
+ # C++ quant kernel (from v25d)
+ # ---------------------------------------------------------------------------
+
HIP_KERNEL = r"""
#include <hip/hip_runtime.h>
#include <stdint.h>
- using int4_v = int __attribute__((ext_vector_type(4)));
- using float4_v = float __attribute__((ext_vector_type(4)));
- using bf16x2 = __bf16 __attribute__((ext_vector_type(2)));
+ using bf16x2 = __bf16 __attribute__((ext_vector_type(2)));
- static constexpr int FP4_E2M1 = 4;
-
- __device__ float4_v __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(
- int4_v a, int4_v b, float4_v c,
- int cbsz, int blgp, int op_sel_a, int scale_a, int op_sel_b, int scale_b
- ) __asm("llvm.amdgcn.mfma.scale.f32.16x16x128.f8f6f4.v4i32.v4i32");
-
- __device__ __forceinline__ int4_v load16(const uint8_t* __restrict__ p) {
- return *reinterpret_cast<const int4_v*>(p);
- }
-
- __device__ __forceinline__ uint16_t float_to_bf16(float f) {
- bf16x2 v;
- v[0] = static_cast<__bf16>(f);
- uint16_t r;
- __builtin_memcpy(&r, &v, sizeof(r));
- return r;
- }
-
__device__ __forceinline__ uint8_t hw_bf16x2_to_fp4x2(uint32_t bf16_pair, float scale) {
uint32_t result;
asm volatile("v_cvt_scalef32_pk_fp4_bf16 %0, %1, %2"
⋯ 13 unchanged lines
const int group = blockIdx.x * 128 + threadIdx.x;
const int row = group / KS;
const int kg = group % KS;
-
if (row >= M) return;
const auto* src = A_bf16 + (long)row * K + kg * 32;
-
float absMax = 1e-10f;
#pragma unroll
for (int i = 0; i < 32; ++i) {
⋯ 7 unchanged lines
A_scale[(long)row * KS + kg] = static_cast<uint8_t>(inv_exp);
const float hw_scale = __builtin_bit_cast(float, static_cast<uint32_t>(inv_exp) << 23);
-
const uint32_t* src_u32 = reinterpret_cast<const uint32_t*>(src);
auto* dst = reinterpret_cast<uint8_t*>(A_fp4 + (long)row * K2 + kg * 16);
#pragma unroll
⋯ 2 unchanged lines
}
}
- template<bool ALWAYS_VALID>
- __device__ __forceinline__ int4_v load_or_zero(bool rt_valid, const uint8_t* p) {
- if constexpr (ALWAYS_VALID) return load16(p);
- else return rt_valid ? load16(p) : int4_v{0,0,0,0};
+ extern "C" void launch_quant(
+ const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale, int M, int K)
+ {
+ const int KS = K / 32;
+ const int n_groups = M * KS;
+ const dim3 block{128};
+ const dim3 grid{static_cast<uint32_t>((n_groups + 127) / 128)};
+ mxfp4_quant<<<grid, block>>>(A_bf16, A_fp4, A_scale, M, K);
}
+ """
- // Main GEMM kernel — fully specialized on NK dimensions
- template<int BM, int BN, int NWARPS, bool SPLITK, bool A_VALID, bool B_VALID,
- int CKT, int C_N, int C_K, int C_SCALEN, int C_TOTAL_KT, int C_KPS>
- __global__ void __launch_bounds__(NWARPS * 64, (NWARPS <= 2) ? 4 : 2)
- mxfp4_gemm(
- const uint8_t* __restrict__ A,
- const uint8_t* __restrict__ As,
- const uint8_t* __restrict__ Bsh,
- const uint8_t* __restrict__ Bssh,
- float* __restrict__ C_partial,
- uint16_t* __restrict__ C_final,
- int M,
- int tile_off_x, int tile_off_y)
- {
- static_assert(BN % 16 == 0);
- constexpr int WAVES_M = (BM + 15) / 16;
- constexpr int WAVES_N = BN / 16;
- static_assert(WAVES_M * WAVES_N == NWARPS);
+ CPP = r"""
+ #include <torch/extension.h>
+ #include <c10/core/DeviceGuard.h>
- constexpr int N = C_N;
- constexpr int K = C_K;
- constexpr int scaleN = C_SCALEN;
- constexpr int K2 = K / 2;
- constexpr int KS = K / 32;
- constexpr long bsh_n_stride = (long)(K / 64) * 512;
+ extern "C" void launch_quant(const __bf16*, uint8_t*, uint8_t*, int, int);
- const int ks_idx = blockIdx.z;
- const int lane = threadIdx.x % 64;
- const int wave = threadIdx.x / 64;
- const int wave_m = wave / WAVES_N;
- const int wave_n = wave % WAVES_N;
+ struct QuantWorkspace {
+ at::Tensor A_fp4;
+ at::Tensor A_scale;
+ int64_t last_M = -1, last_K = -1;
+ void ensure(int M, int K, const at::TensorOptions& opts) {
+ if (M == last_M && K == last_K) return;
+ A_fp4 = at::empty({(int64_t)M, (int64_t)(K / 2)}, opts.dtype(at::kByte));
+ A_scale = at::empty({(int64_t)M, (int64_t)(K / 32)}, opts.dtype(at::kByte));
+ last_M = M; last_K = K;
+ }
+ };
- const int tile_m = (blockIdx.y + tile_off_y) * BM + wave_m * 16;
- const int tile_n = (blockIdx.x + tile_off_x) * BN + wave_n * 16;
+ static QuantWorkspace g_qws;
- if (tile_m >= M || tile_n >= N) return;
+ std::vector<at::Tensor> do_quant(const at::Tensor& A) {
+ auto guard = at::DeviceGuard(A.device());
+ at::Tensor A_bf16 = (A.scalar_type() == at::kBFloat16 && A.is_contiguous())
+ ? A : A.to(A.device(), at::kBFloat16, false, false,
+ at::MemoryFormat::Contiguous);
+ const int M = A_bf16.size(0);
+ const int K = A_bf16.size(1);
+ g_qws.ensure(M, K, A_bf16.options());
+ launch_quant(
+ reinterpret_cast<const __bf16*>(A_bf16.data_ptr<at::BFloat16>()),
+ g_qws.A_fp4.data_ptr<uint8_t>(),
+ g_qws.A_scale.data_ptr<uint8_t>(),
+ M, K);
+ return {g_qws.A_fp4, g_qws.A_scale};
+ }
+ """
- constexpr int ktiles_per_split = C_KPS;
- const int ks_start = ks_idx * ktiles_per_split;
+ _ext = load_inline(
+ name=f"g_{uuid.uuid4().hex[:8]}",
+ cpp_sources=[CPP],
+ cuda_sources=[HIP_KERNEL],
+ functions=["do_quant"],
+ with_cuda=True,
+ extra_cflags=["-O3", "-std=c++20"],
+ extra_cuda_cflags=[
+ "-O3", "--offload-arch=gfx950", "-ffast-math", "-munsafe-fp-atomics",
+ "-std=c++20", "-mllvm", "-amdgpu-early-inline-all=true",
+ "-mllvm", "-amdgpu-function-calls=false", "-mwavefrontsize64",
+ "-mcumode", "-fgpu-flush-denormals-to-zero",
+ ],
+ extra_ldflags=["-lamdhip64"],
+ )
- const int lrow = lane % 16;
- const int kgrp = lane / 16;
- const int gm = tile_m + lrow;
- const int gn = tile_n + lrow;
+ # ---------------------------------------------------------------------------
+ # Triton helpers
+ # ---------------------------------------------------------------------------
- const bool a_rt = A_VALID | (gm < M);
- const bool b_rt = B_VALID | (gn < N);
+ @triton.jit
+ def remap_xcd(pid, GRID_MN, NUM_XCDS: tl.constexpr = 8):
+ pids_per_xcd = (GRID_MN + NUM_XCDS - 1) // NUM_XCDS
+ tall_xcds = GRID_MN % NUM_XCDS
+ tall_xcds = NUM_XCDS if tall_xcds == 0 else tall_xcds
+ xcd = pid % NUM_XCDS
+ local_pid = pid // NUM_XCDS
+ if xcd < tall_xcds:
+ pid = xcd * pids_per_xcd + local_pid
+ else:
+ pid = tall_xcds * pids_per_xcd + (xcd - tall_xcds) * (pids_per_xcd - 1) + local_pid
+ return pid
- const int n_tile = tile_n / 16;
- const auto bsh_lane_base = Bsh + (long)n_tile * bsh_n_stride + (long)lrow * 16;
- const uint8_t* a_row = nullptr;
- const uint8_t* as_row = nullptr;
- if constexpr (A_VALID) {
- a_row = A + (long)gm * K2;
- as_row = As + (long)gm * KS;
- } else {
- if (a_rt) { a_row = A + (long)gm * K2;
- as_row = As + (long)gm * KS; }
- }
+ @triton.jit
+ def pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M: tl.constexpr = 1):
+ if GROUP_SIZE_M == 1:
+ pid_m = pid // num_pid_n
+ pid_n = pid % num_pid_n
+ else:
+ num_pid_in_group = GROUP_SIZE_M * num_pid_n
+ group_id = pid // num_pid_in_group
+ first_pid_m = group_id * GROUP_SIZE_M
+ group_size_m = tl.minimum(num_pid_m - first_pid_m, GROUP_SIZE_M)
+ pid_m = first_pid_m + (pid % group_size_m)
+ pid_n = (pid % num_pid_in_group) // group_size_m
+ return pid_m, pid_n
- int bssh_base = 0;
- if constexpr (B_VALID) {
- bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64 + (gn >> 5) * (32 * scaleN);
- } else {
- if (b_rt) bssh_base = ((gn >> 4) & 1) + (gn & 15) * 4 + kgrp * 64 + (gn >> 5) * (32 * scaleN);
- }
- const int k_half_off = (kgrp & 1) * 256;
- const int k_blk_base = kgrp >> 1;
+ # ---------------------------------------------------------------------------
+ # Triton GEMM kernel (identical to v25t)
+ # ---------------------------------------------------------------------------
- static constexpr long a_kt_stride = 64L;
- static constexpr long bsh_kt_stride = 1024L;
+ @triton.heuristics({
+ "EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)
+ and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
+ and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),
+ })
+ @triton.jit
+ def _mxfp4_gemm_kernel(
+ a_ptr, b_ptr, c_ptr, a_scales_ptr, b_scales_ptr,
+ M, N, K,
+ stride_am, stride_ak, stride_bn, stride_bk,
+ stride_ck, stride_cm, stride_cn,
+ stride_asm, stride_ask, stride_bsn, stride_bsk,
+ BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
+ BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr,
+ NUM_KSPLIT: tl.constexpr, SPLITK_BLOCK_SIZE: tl.constexpr,
+ EVEN_K: tl.constexpr, num_warps: tl.constexpr,
+ num_stages: tl.constexpr, waves_per_eu: tl.constexpr,
+ matrix_instr_nonkdim: tl.constexpr, cache_modifier: tl.constexpr,
+ ):
+ tl.assume(stride_am > 0); tl.assume(stride_ak > 0)
+ tl.assume(stride_bk > 0); tl.assume(stride_bn > 0)
+ tl.assume(stride_cm > 0); tl.assume(stride_cn > 0)
+ tl.assume(stride_asm > 0); tl.assume(stride_ask > 0)
+ tl.assume(stride_bsk > 0); tl.assume(stride_bsn > 0)
- const uint8_t* a_ptr = nullptr;
- const uint8_t* bsh_ptr = nullptr;
- const uint8_t* bssh_ptr = nullptr;
+ GRID_MN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)
+ SCALE_GROUP_SIZE: tl.constexpr = 32
- if constexpr (A_VALID) {
- a_ptr = a_row + (long)ks_start * a_kt_stride + kgrp * 16;
- } else {
- if (a_rt) a_ptr = a_row + (long)ks_start * a_kt_stride + kgrp * 16;
- }
- if constexpr (B_VALID) {
- bsh_ptr = bsh_lane_base + (long)(ks_start * 2 + k_blk_base) * 512 + k_half_off;
- bssh_ptr = Bssh + bssh_base + (ks_start & 1) * 2 + (ks_start >> 1) * 256;
- } else {
- if (b_rt) {
- bsh_ptr = bsh_lane_base + (long)(ks_start * 2 + k_blk_base) * 512 + k_half_off;
- bssh_ptr = Bssh + bssh_base + (ks_start & 1) * 2 + (ks_start >> 1) * 256;
- }
- }
+ pid_unified = tl.program_id(axis=0)
+ pid_unified = remap_xcd(pid_unified, GRID_MN * NUM_KSPLIT, NUM_XCDS=8)
+ pid_k = pid_unified % NUM_KSPLIT
+ pid = pid_unified // NUM_KSPLIT
+ num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
+ num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
- float4_v acc{0.f, 0.f, 0.f, 0.f};
+ if NUM_KSPLIT == 1:
+ pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
+ else:
+ pid_m = pid // num_pid_n
+ pid_n = pid % num_pid_n
- const int bssh_step0 = (ks_start & 1) ? 254 : 2;
- const int bssh_step1 = 256 - bssh_step0;
+ tl.assume(pid_m >= 0); tl.assume(pid_n >= 0)
- #define DO_MFMA(a_off, b_off, bssh_off, ks_val) \
- { \
- const auto av = load_or_zero<A_VALID>(a_rt, a_ptr + (a_off) * a_kt_stride); \
- const auto bv = load_or_zero<B_VALID>(b_rt, bsh_ptr + (b_off) * bsh_kt_stride); \
- const int ks = (ks_val) * 4 + kgrp; \
- int sa, sb; \
- if constexpr (A_VALID) sa = static_cast<int>(as_row[ks]); \
- else sa = (a_rt & (ks < KS)) ? static_cast<int>(as_row[ks]) : 127; \
- if constexpr (B_VALID) sb = static_cast<int>(*(bssh_ptr + (bssh_off))); \
- else sb = (b_rt & (ks < KS)) ? static_cast<int>(*(bssh_ptr + (bssh_off))) : 127; \
- acc = __builtin_amdgcn_mfma_scale_f32_16x16x128_f8f6f4(av,bv,acc,FP4_E2M1,FP4_E2M1,0,sa,0,sb); \
- }
+ if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
+ num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
- static_assert(CKT > 0, "Specialized kernel must have compile-time CKT");
- #pragma unroll
- for (int q = 0; q < (CKT / 4); ++q) {
- DO_MFMA(0, 0, 0, ks_start + q*4)
- DO_MFMA(1, 1, bssh_step0, ks_start + q*4 + 1)
- DO_MFMA(2, 2, bssh_step0 + bssh_step1, ks_start + q*4 + 2)
- DO_MFMA(3, 3, bssh_step0 + bssh_step1 + bssh_step0, ks_start + q*4 + 3)
- a_ptr += 4 * a_kt_stride;
- bsh_ptr += 4 * bsh_kt_stride;
- bssh_ptr += 512;
- }
- if constexpr ((CKT % 4) >= 2) {
- DO_MFMA(0, 0, 0, ks_start + (CKT/4)*4)
- DO_MFMA(1, 1, bssh_step0, ks_start + (CKT/4)*4 + 1)
- a_ptr += 2 * a_kt_stride;
- bsh_ptr += 2 * bsh_kt_stride;
- bssh_ptr += 256;
- }
- if constexpr ((CKT % 2) == 1) {
- DO_MFMA(0, 0, 0, ks_start + CKT - 1)
- }
- #undef DO_MFMA
+ offs_k = tl.arange(0, BLOCK_SIZE_K // 2)
+ offs_k_split = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_k
+ offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
+ a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_k_split[None, :] * stride_ak
- const int out_col = tile_n + lrow;
- const int out_row_base = tile_m + kgrp * 4;
+ offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
+ offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
+ offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
+ b_ptrs = b_ptr + offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk
- if constexpr (!B_VALID) { if (out_col >= N) return; }
+ offs_ks_a = pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) + tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
+ a_scale_ptrs = a_scales_ptr + offs_am[:, None] * stride_asm + offs_ks_a[None, :] * stride_ask
- constexpr bool out_rows_always_valid = A_VALID && (BM >= 16);
+ offs_asn = (pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)) % N
+ offs_ks_b = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)
+ b_scale_ptrs = b_scales_ptr + offs_asn[:, None] * stride_bsn + offs_ks_b[None, :] * stride_bsk
- if constexpr (SPLITK) {
- auto c_out = C_partial + (long)ks_idx * M * N + (long)out_row_base * N + out_col;
- #pragma unroll
- for (int i = 0; i < 4; ++i) {
- if constexpr (out_rows_always_valid) c_out[i * N] = acc[i];
- else if (out_row_base + i < M) c_out[i * N] = acc[i];
- }
- } else {
- auto c_out = C_final + (long)out_row_base * N + out_col;
- #pragma unroll
- for (int i = 0; i < 4; ++i) {
- if constexpr (out_rows_always_valid) c_out[i * N] = float_to_bf16(acc[i]);
- else if (out_row_base + i < M) c_out[i * N] = float_to_bf16(acc[i]);
- }
+ accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
+
+ for k in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
+ a_scales = tl.load(a_scale_ptrs)
+ b_scales = (
+ tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
+ .reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8, 4, 16, 2, 2, 1)
+ .permute(0, 5, 3, 1, 4, 2, 6)
+ .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
+ )
+ if EVEN_K:
+ a = tl.load(a_ptrs)
+ b = tl.load(b_ptrs, cache_modifier=cache_modifier)
+ else:
+ a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k * (BLOCK_SIZE_K // 2), other=0)
+ b = tl.load(b_ptrs, mask=offs_k_shuffle_arr[None, :] < (K - k * (BLOCK_SIZE_K // 2)) * 16, other=0, cache_modifier=cache_modifier)
+
+ b = b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16).permute(0, 1, 4, 2, 3, 5).reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2).trans(1, 0)
+ accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)
+
+ a_ptrs += (BLOCK_SIZE_K // 2) * stride_ak
+ b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
+ a_scale_ptrs += (BLOCK_SIZE_K // SCALE_GROUP_SIZE) * stride_ask
+ b_scale_ptrs += BLOCK_SIZE_K * stride_bsk
+
+ c = accumulator.to(c_ptr.type.element_ty)
+ offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
+ offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
+ c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :] + pid_k * stride_ck
+ c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
+ tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")
+
+
+ @triton.jit
+ def _mxfp4_reduce_kernel(
+ c_in_ptr, c_out_ptr, M, N,
+ stride_c_in_k, stride_c_in_m, stride_c_in_n,
+ stride_c_out_m, stride_c_out_n,
+ BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
+ ACTUAL_KSPLIT: tl.constexpr, MAX_KSPLIT: tl.constexpr,
+ ):
+ pid_m = tl.program_id(axis=0)
+ pid_n = tl.program_id(axis=1)
+ offs_m = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
+ offs_n = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
+ offs_k = tl.arange(0, MAX_KSPLIT)
+ c_in_ptrs = c_in_ptr + offs_k[:, None, None] * stride_c_in_k + offs_m[None, :, None] * stride_c_in_m + offs_n[None, None, :] * stride_c_in_n
+ if ACTUAL_KSPLIT == MAX_KSPLIT:
+ c = tl.load(c_in_ptrs)
+ else:
+ c = tl.load(c_in_ptrs, mask=offs_k[:, None, None] < ACTUAL_KSPLIT)
+ c = tl.sum(c, axis=0).to(c_out_ptr.type.element_ty)
+ c_out_ptrs = c_out_ptr + offs_m[:, None] * stride_c_out_m + offs_n[None, :] * stride_c_out_n
+ tl.store(c_out_ptrs, c)
+
+
+ # ---------------------------------------------------------------------------
+ # AITER default configs (fallback / seeds)
+ # ---------------------------------------------------------------------------
+
+ def _cfg(bm, bn, bk, gsm, nw, ns, wpe, cm, nks):
+ return {
+ "BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": bn, "BLOCK_SIZE_K": bk,
+ "GROUP_SIZE_M": gsm, "num_warps": nw, "num_stages": ns,
+ "waves_per_eu": wpe, "matrix_instr_nonkdim": 16,
+ "cache_modifier": cm, "NUM_KSPLIT": nks,
}
+
+ _DEFAULTS = {
+ (4, 2880, 512): _cfg(8, 64, 512, 1, 2, 1, 1, None, 1),
+ (16, 2112, 7168): _cfg(16, 32, 512, 1, 4, 1, 4, ".cg", 14),
+ (32, 4096, 512): _cfg(32, 64, 512, 1, 4, 1, 1, None, 1),
+ (32, 2880, 512): _cfg(32, 64, 512, 1, 4, 1, 1, None, 1),
+ (64, 7168, 2048): _cfg(32, 64, 512, 4, 2, 2, 1, None, 1),
+ (256,3072, 1536): _cfg(128,32, 512, 1, 4, 2, 1, None, 1),
+ # test shapes
+ (8, 2112, 7168): _cfg(8, 32, 512, 1, 2, 2, 1, ".cg", 7),
+ (16, 3072, 1536): _cfg(8, 32, 512, 1, 4, 2, 1, None, 1),
+ (64, 3072, 1536): _cfg(64, 32, 512, 1, 2, 2, 1, ".cg", 1),
+ (256,2880, 512): _cfg(32, 64, 512, 1, 2, 2, 1, None, 1),
}
- template<int C_N>
- __global__ void mxfp4_reduce(
- const float* __restrict__ C_partial,
- uint16_t* __restrict__ C_out,
- int M, int NUM_KSPLIT)
- {
- constexpr int N = C_N;
- const int col = blockIdx.x * 32 + threadIdx.x;
- const int row = blockIdx.y * 16 + threadIdx.y;
- if (row >= M || col >= N) return;
- float sum = 0.f;
- const long mn = (long)row * N + col;
- const long mn_stride = (long)M * N;
- for (int k = 0; k < NUM_KSPLIT; ++k)
- sum += C_partial[k * mn_stride + mn];
+ # ---------------------------------------------------------------------------
+ # splitK helper (from AITER)
+ # ---------------------------------------------------------------------------
- bf16x2 v;
- v[0] = static_cast<__bf16>(sum);
- uint16_t r;
- __builtin_memcpy(&r, &v, sizeof(r));
- C_out[mn] = r;
- }
+ 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
+ NUM_KSPLIT = triton.cdiv(K, SPLITK_BLOCK_SIZE // 2)
+ return SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT
- extern "C" void launch_quant(
- const __bf16* A_bf16, uint8_t* A_fp4, uint8_t* A_scale, int M, int K)
- {
- const int KS = K / 32;
- const int n_groups = M * KS;
- const dim3 block{128};
- const dim3 grid{static_cast<uint32_t>((n_groups + 127) / 128)};
- mxfp4_quant<<<grid, block>>>(A_bf16, A_fp4, A_scale, M, K);
- }
- template<int C_N, int C_K, int C_SCALEN, int C_TOTAL_KT, int C_KPS, int C_NUM_KSPLIT>
- void launch_gemm_nk(
- const uint8_t* A, const uint8_t* As,
- const uint8_t* Bsh, const uint8_t* Bssh,
- float* C_partial, uint16_t* C_final, int M)
- {
- constexpr bool do_splitk = C_NUM_KSPLIT > 1;
+ # ---------------------------------------------------------------------------
+ # Kernel runner: launches GEMM (+ reduce) with a given config
+ # ---------------------------------------------------------------------------
- auto launch = [&]<int BM, int BN, int NWARPS>() {
- static_assert(((BM + 15) / 16) * (BN / 16) == NWARPS);
+ def _run_gemm(config, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed, y, y_pp):
+ cfg = dict(config)
+ BK = cfg["BLOCK_SIZE_K"]
+ NUM_KSPLIT = cfg["NUM_KSPLIT"]
- const int full_m = (BM >= 16) ? M / BM : 0;
- constexpr int full_n = C_N / BN;
- const int total_m = (M + BM - 1) / BM;
- constexpr int total_n = (C_N + BN - 1) / BN;
- const int edge_m = total_m - full_m;
- constexpr int edge_n = total_n - full_n;
+ if BK >= 2 * K_packed:
+ BK = triton.next_power_of_2(2 * K_packed)
+ cfg["BLOCK_SIZE_K"] = BK
+ cfg["NUM_KSPLIT"] = 1
+ NUM_KSPLIT = 1
- const dim3 block{static_cast<uint32_t>(NWARPS * 64)};
+ cfg["BLOCK_SIZE_K"] = max(cfg["BLOCK_SIZE_K"], 256)
+ BK = cfg["BLOCK_SIZE_K"]
- auto sub = [&]<bool AV, bool BV>(int gx, int gy, int ox, int oy) {
- if (gx <= 0 || gy <= 0) return;
- const dim3 grid{
- static_cast<uint32_t>(gx),
- static_cast<uint32_t>(gy),
- static_cast<uint32_t>(C_NUM_KSPLIT)
- };
- if constexpr (do_splitk)
- mxfp4_gemm<BM,BN,NWARPS,true,AV,BV,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS>
- <<<grid,block>>>(A,As,Bsh,Bssh,C_partial,nullptr,M,ox,oy);
- else
- mxfp4_gemm<BM,BN,NWARPS,false,AV,BV,C_KPS,C_N,C_K,C_SCALEN,C_TOTAL_KT,C_KPS>
- <<<grid,block>>>(A,As,Bsh,Bssh,nullptr,C_final,M,ox,oy);
- };
+ if NUM_KSPLIT > 1:
+ SPLITK_BS, BK, NUM_KSPLIT = get_splitk(K_packed, BK, NUM_KSPLIT)
+ cfg["BLOCK_SIZE_K"] = BK
+ cfg["NUM_KSPLIT"] = NUM_KSPLIT
+ cfg["SPLITK_BLOCK_SIZE"] = SPLITK_BS
+ else:
+ cfg["SPLITK_BLOCK_SIZE"] = 2 * K_packed
- sub.template operator()<true, true >(full_n, full_m, 0, 0);
- sub.template operator()<true, false>(edge_n, full_m, full_n, 0);
- sub.template operator()<false, true >(full_n, edge_m, 0, full_m);
- sub.template operator()<false, false>(edge_n, edge_m, full_n, full_m);
+ c_out = y_pp if NUM_KSPLIT > 1 else y
- if constexpr (do_splitk) {
- const dim3 rblock{32, 16};
- const dim3 rgrid{
- static_cast<uint32_t>((C_N + 31) / 32),
- static_cast<uint32_t>((M + 15) / 16)
- };
- mxfp4_reduce<C_N><<<rgrid, rblock>>>(C_partial, C_final, M, C_NUM_KSPLIT);
- }
- };
+ KS = K_elem // 32
+ scaleN = ((KS + 7) // 8) * 8
- if (M <= 8) launch.template operator()< 8, 32, 2>();
- else if (M <= 16) launch.template operator()< 16, 32, 2>();
- else if (M <= 32) launch.template operator()< 16, 32, 2>();
- else if (M <= 64) launch.template operator()< 32, 32, 4>();
- else if (M <=128) launch.template operator()< 32, 32, 4>();
- else launch.template operator()< 64, 32, 8>();
- }
+ grid = lambda META: (
+ META["NUM_KSPLIT"]
+ * triton.cdiv(M, META["BLOCK_SIZE_M"])
+ * triton.cdiv(N, META["BLOCK_SIZE_N"]),
+ )
- template<int C_N, int C_K, int C_SCALEN, int C_TOTAL_KT>
- void launch_gemm_nk_7168(
- const uint8_t* A, const uint8_t* As,
- const uint8_t* Bsh, const uint8_t* Bssh,
- float* C_partial, uint16_t* C_final, int M)
- {
- if (M <= 8)
- launch_gemm_nk<C_N, C_K, C_SCALEN, C_TOTAL_KT, 8, 7>(A, As, Bsh, Bssh, C_partial, C_final, M);
- else if (M <= 16)
- launch_gemm_nk<C_N, C_K, C_SCALEN, C_TOTAL_KT, 4, 14>(A, As, Bsh, Bssh, C_partial, C_final, M);
- else
- launch_gemm_nk<C_N, C_K, C_SCALEN, C_TOTAL_KT, 56, 1>(A, As, Bsh, Bssh, C_partial, C_final, M);
- }
+ _mxfp4_gemm_kernel[grid](
+ A_fp4, w, c_out, A_scale, b_scales,
+ M, N, K_packed,
+ A_fp4.stride(0), A_fp4.stride(1),
+ w.stride(0), w.stride(1),
+ 0 if NUM_KSPLIT == 1 else y_pp.stride(0),
+ c_out.stride(-2), c_out.stride(-1),
+ A_scale.stride(0), A_scale.stride(1),
+ 32 * scaleN, 1,
+ **cfg,
+ )
- // Precomputed raw-pointer fast path — avoids ALL tensor ops in hot path
- extern "C" void launch_gemm_raw(
- const uint8_t* A_fp4, const uint8_t* A_scale,
- const uint8_t* Bsh, const uint8_t* Bssh,
- float* C_partial, uint16_t* C_final,
- int M, int N, int K)
- {
- if (N == 2880 && K == 512)
- launch_gemm_nk<2880, 512, 16, 4, 4, 1>(A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
- else if (N == 2112 && K == 7168)
- launch_gemm_nk_7168<2112, 7168, 224, 56>(A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
- else if (N == 4096 && K == 512)
- launch_gemm_nk<4096, 512, 16, 4, 4, 1>(A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
- else if (N == 7168 && K == 2048)
- launch_gemm_nk<7168, 2048, 64, 16, 16, 1>(A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
- else if (N == 3072 && K == 1536)
- launch_gemm_nk<3072, 1536, 48, 12, 12, 1>(A_fp4, A_scale, Bsh, Bssh, C_partial, C_final, M);
- }
- """
+ if NUM_KSPLIT > 1:
+ ACTUAL_KSPLIT = triton.cdiv(K_packed, cfg["SPLITK_BLOCK_SIZE"] // 2)
+ grid_r = (triton.cdiv(M, 16), triton.cdiv(N, 64))
+ _mxfp4_reduce_kernel[grid_r](
+ y_pp, y, M, N,
+ y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
+ y.stride(0), y.stride(1),
+ 16, 64, ACTUAL_KSPLIT, triton.next_power_of_2(NUM_KSPLIT),
+ )
+ return y
- CPP = r"""
- #include <torch/extension.h>
- #include <c10/core/DeviceGuard.h>
- extern "C" void launch_quant(const __bf16*, uint8_t*, uint8_t*, int, int);
- extern "C" void launch_gemm_raw(const uint8_t*, const uint8_t*, const uint8_t*, const uint8_t*,
- float*, uint16_t*, int, int, int);
+ # ---------------------------------------------------------------------------
+ # Two-phase exhaustive config search
+ # ---------------------------------------------------------------------------
- static int get_num_ksplit(int M, int K) {
- int total_ktiles = K / 128;
- if (total_ktiles < 28) return 1;
- if (M <= 8) return 7;
- if (M <= 16) return 14;
- return 1;
- }
+ _CONFIG_CACHE = {}
- // Precomputed workspace — keyed by (M, N, K) tuple
- struct ShapeWorkspace {
- at::Tensor A_fp4;
- at::Tensor A_scale;
- at::Tensor C_partial;
- at::Tensor C;
- uint8_t* a_fp4_ptr = nullptr;
- uint8_t* a_scale_ptr = nullptr;
- float* c_partial_ptr = nullptr;
- uint16_t* c_final_ptr = nullptr;
- int M = 0, N = 0, K = 0;
- int num_ksplit = 0;
- };
- // Cache for B tensor pointers (B doesn't change between calls for same N,K)
- struct BCache {
- const uint8_t* bsh_ptr = nullptr;
- const uint8_t* bssh_ptr = nullptr;
- int64_t bsh_data_ptr = 0; // for staleness check
- int64_t bssh_data_ptr = 0;
- };
+ def _validate_config(config, K_packed):
+ BN = config["BLOCK_SIZE_N"]
+ BK = config["BLOCK_SIZE_K"]
+ if BN < 32:
+ return False
+ if BK < 64:
+ return False
+ if K_packed % (BK // 2) != 0:
+ return False
+ return True
- // Up to 10 different (M,N,K) combos (4 test + 6 bench)
- static ShapeWorkspace g_ws[10];
- static int g_ws_count = 0;
- static BCache g_bcache;
- static ShapeWorkspace* find_or_create_ws(int M, int N, int K, int num_ksplit,
- const at::TensorOptions& opts) {
- // Search existing
- for (int i = 0; i < g_ws_count; ++i) {
- if (g_ws[i].M == M && g_ws[i].N == N && g_ws[i].K == K)
- return &g_ws[i];
- }
- // Create new
- auto& ws = g_ws[g_ws_count++];
- ws.M = M; ws.N = N; ws.K = K;
- ws.num_ksplit = num_ksplit;
- int64_t KS = K / 32;
- ws.A_fp4 = at::empty({(int64_t)M, (int64_t)(K / 2)}, opts.dtype(at::kByte));
- ws.A_scale = at::empty({(int64_t)M, KS}, opts.dtype(at::kByte));
- ws.C = at::empty({(int64_t)M, (int64_t)N}, opts.dtype(at::kBFloat16));
- if (num_ksplit > 1)
- ws.C_partial = at::empty({(int64_t)num_ksplit, (int64_t)M, (int64_t)N}, opts.dtype(at::kFloat));
- // Cache raw pointers
- ws.a_fp4_ptr = ws.A_fp4.data_ptr<uint8_t>();
- ws.a_scale_ptr = ws.A_scale.data_ptr<uint8_t>();
- ws.c_partial_ptr = (num_ksplit > 1) ? ws.C_partial.data_ptr<float>() : nullptr;
- ws.c_final_ptr = reinterpret_cast<uint16_t*>(ws.C.data_ptr<at::BFloat16>());
- return &ws;
- }
+ def _time_config(config, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed,
+ n_warmup=2, n_iter=8):
+ NUM_KSPLIT = config.get("NUM_KSPLIT", 1)
+ y = torch.empty((M, N), dtype=torch.bfloat16, device=A_fp4.device)
+ y_pp = None
+ if NUM_KSPLIT > 1:
+ y_pp = torch.empty((NUM_KSPLIT, M, N), dtype=torch.float32, device=A_fp4.device)
- at::Tensor fwd(const at::Tensor& A,
- const at::Tensor& B_q,
- const at::Tensor& B_shuffle,
- const at::Tensor& B_scale_sh) {
- auto guard = at::DeviceGuard(A.device());
+ for _ in range(n_warmup):
+ _run_gemm(config, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed, y, y_pp)
- const int M = A.size(0);
- const int K = A.size(1);
- const int N = B_q.size(0);
+ torch.cuda.synchronize()
+ start_evt = torch.cuda.Event(enable_timing=True)
+ end_evt = torch.cuda.Event(enable_timing=True)
+ start_evt.record()
+ for _ in range(n_iter):
+ _run_gemm(config, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed, y, y_pp)
+ end_evt.record()
+ end_evt.synchronize()
- // Fast path: check if A is already bf16 contiguous
- const __bf16* a_bf16_ptr;
- at::Tensor A_bf16;
- if (A.scalar_type() == at::kBFloat16 && A.is_contiguous()) {
- a_bf16_ptr = reinterpret_cast<const __bf16*>(A.data_ptr<at::BFloat16>());
- } else {
- A_bf16 = A.to(A.device(), at::kBFloat16, false, false, at::MemoryFormat::Contiguous);
- a_bf16_ptr = reinterpret_cast<const __bf16*>(A_bf16.data_ptr<at::BFloat16>());
- }
+ return start_evt.elapsed_time(end_evt) / n_iter * 1000 # ms -> us
- // Cache B pointers — B tensors don't change between benchmark iterations
- auto bsh_dp = reinterpret_cast<int64_t>(B_shuffle.data_ptr());
- auto bssh_dp = reinterpret_cast<int64_t>(B_scale_sh.data_ptr());
- if (bsh_dp != g_bcache.bsh_data_ptr || bssh_dp != g_bcache.bssh_data_ptr) {
- // First call or B changed — resolve views once
- at::Tensor Bsh = B_shuffle.view(at::kByte);
- if (!Bsh.is_contiguous()) Bsh = Bsh.contiguous();
- at::Tensor Bssh = B_scale_sh.view(at::kByte);
- if (!Bssh.is_contiguous()) Bssh = Bssh.contiguous();
- g_bcache.bsh_ptr = Bsh.data_ptr<uint8_t>();
- g_bcache.bssh_ptr = Bssh.data_ptr<uint8_t>();
- g_bcache.bsh_data_ptr = bsh_dp;
- g_bcache.bssh_data_ptr = bssh_dp;
- }
- const int num_ksplit = get_num_ksplit(M, K);
- auto* ws = find_or_create_ws(M, N, K, num_ksplit, A.options());
+ def _get_param_choices(M, N, K_elem):
+ K_packed = K_elem // 2
- // Quant: A_bf16 -> A_fp4 + A_scale
- launch_quant(a_bf16_ptr, ws->a_fp4_ptr, ws->a_scale_ptr, M, K);
+ bm_choices = [b for b in [8, 16, 32, 64, 128, 256] if b <= max(M * 2, 8)]
+ bn_choices = [b for b in [32, 64, 128, 256] if b <= N]
+ bk_choices = [b for b in [256, 512, 1024] if K_packed % (b // 2) == 0]
+ if not bk_choices:
+ bk_choices = [256]
- // GEMM: all raw pointers, no tensor ops
- launch_gemm_raw(
- ws->a_fp4_ptr, ws->a_scale_ptr,
- g_bcache.bsh_ptr, g_bcache.bssh_ptr,
- ws->c_partial_ptr, ws->c_final_ptr,
- M, N, K);
+ max_splits = K_packed // (min(bk_choices) // 2)
+ sk_choices = [1]
+ for s in [2, 3, 4, 7, 14]:
+ if s <= max_splits and K_packed % s == 0:
+ sk_choices.append(s)
- return ws->C;
- }
- """
+ return bm_choices, bn_choices, bk_choices, sk_choices
- _ext = load_inline(
- name=f"g_{uuid.uuid4().hex[:8]}",
- cpp_sources=[CPP],
- cuda_sources=[HIP_KERNEL],
- functions=["fwd"],
- with_cuda=True,
- extra_cflags=["-O3", "-std=c++20"],
- extra_cuda_cflags=[
- "-O3",
- "--offload-arch=gfx950",
- "-ffast-math",
- "-ffinite-math-only",
- "-munsafe-fp-atomics",
- "-std=c++20",
- "-mllvm", "-amdgpu-early-inline-all=true",
- "-mllvm", "-amdgpu-function-calls=false",
- "-mwavefrontsize64",
- "-mcumode",
- "-mllvm", "--amdgpu-kernarg-preload-count=16",
- "-mllvm", "-enable-post-misched=0",
- "-mllvm", "--lsr-drop-solution=1",
- "-mllvm", "-amdgpu-coerce-illegal-types=1",
- "-fgpu-flush-denormals-to-zero",
- "-fno-offload-uniform-block",
- # New aggressive flags
- "-mllvm", "-amdgpu-loop-prefetch=true",
- "-mllvm", "-enable-unroll-and-jam=true",
- "-mllvm", "-amdgpu-set-wave-priority=true",
- "-mllvm", "-unroll-threshold=1000",
- "-mllvm", "-amdgpu-internalize-symbols=true",
- ],
- extra_ldflags=["-lamdhip64"],
- )
+ def _search_config(M, N, K_elem, A_fp4, A_scale, w, b_scales):
+ K_packed = K_elem // 2
+ bm_choices, bn_choices, bk_choices, sk_choices = _get_param_choices(M, N, K_elem)
+ # Phase 1: sweep tile params (BM, BN, BK, NUM_KSPLIT) with sensible defaults
+ tile_combos = list(itertools.product(bm_choices, bn_choices, bk_choices, sk_choices))
+ p1_total = len(tile_combos)
+ print(f" Phase 1: {p1_total} tile combos (BM x BN x BK x SK = "
+ f"{len(bm_choices)}x{len(bn_choices)}x{len(bk_choices)}x{len(sk_choices)})", flush=True)
+
+ best_time = float("inf")
+ best_tile = None
+ best_config = _DEFAULTS.get((M, N, K_elem), _cfg(16, 64, 256, 4, 2, 2, 1, None, 1))
+
+ for i, (bm, bn, bk, sk) in enumerate(tile_combos):
+ 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": 1, "matrix_instr_nonkdim": 16,
+ "cache_modifier": None, "NUM_KSPLIT": sk,
+ }
+ if not _validate_config(cfg, K_packed):
+ continue
+ try:
+ t = _time_config(cfg, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed)
+ except Exception:
+ continue
+ if t < best_time:
+ best_time = t
+ best_tile = (bm, bn, bk, sk)
+ best_config = cfg
+ print(f" [P1 {i+1}/{p1_total}] best={t:.1f}us BM={bm} BN={bn} BK={bk} SK={sk}", flush=True)
+
+ if best_tile is None:
+ print(" Phase 1: no valid tile found, using default", flush=True)
+ return best_config
+
+ bm, bn, bk, sk = best_tile
+ print(f" Phase 1 winner: BM={bm} BN={bn} BK={bk} SK={sk} = {best_time:.1f}us", flush=True)
+
+ # Phase 2: sweep (num_warps, num_stages, waves_per_eu, GROUP_SIZE_M, cache_modifier)
+ tune_combos = list(itertools.product(
+ [2, 4], # num_warps
+ [1, 2], # num_stages
+ [1, 2, 4], # waves_per_eu
+ [1, 4, 8], # GROUP_SIZE_M
+ [None, ".cg"], # cache_modifier
+ ))
+ p2_total = len(tune_combos)
+ print(f" Phase 2: {p2_total} tune combos around best tile", flush=True)
+
+ for i, (nw, ns, wpe, gsm, cm) in enumerate(tune_combos):
+ cfg = {
+ "BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": bn, "BLOCK_SIZE_K": bk,
+ "GROUP_SIZE_M": gsm, "num_warps": nw, "num_stages": ns,
+ "waves_per_eu": wpe, "matrix_instr_nonkdim": 16,
+ "cache_modifier": cm, "NUM_KSPLIT": sk,
+ }
+ try:
+ t = _time_config(cfg, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed)
+ except Exception:
+ continue
+ if t < best_time:
+ best_time = t
+ best_config = cfg
+ print(f" [P2 {i+1}/{p2_total}] best={t:.1f}us nw={nw} ns={ns} wpe={wpe} gsm={gsm} cm={cm}", flush=True)
+
+ print(f" Phase 2 winner: {best_time:.1f}us", flush=True)
+ return best_config
+
+
+ # ---------------------------------------------------------------------------
+ # Workspace caching
+ # ---------------------------------------------------------------------------
+
+ class _Workspace:
+ __slots__ = ["y", "y_pp", "_key"]
+ def __init__(self):
+ self.y = None; self.y_pp = None; self._key = None
+ def ensure(self, M, N, num_ksplit, device):
+ key = (M, N, num_ksplit)
+ if key == self._key: return
+ self._key = key
+ self.y = torch.empty((M, N), dtype=torch.bfloat16, device=device)
+ if num_ksplit > 1:
+ self.y_pp = torch.empty((num_ksplit, M, N), dtype=torch.float32, device=device)
+ else:
+ self.y_pp = None
+
+ _ws = _Workspace()
+ _b_cache = {"bsh_dp": 0, "bssh_dp": 0, "w": None, "b_scales": None}
+
+
+ # ---------------------------------------------------------------------------
+ # Main entry point
+ # ---------------------------------------------------------------------------
+
def custom_kernel(data: Tuple[torch.Tensor, ...]) -> torch.Tensor:
- """MXFP4 GEMM v25d: v25b + aggressive compiler flags."""
+ """MXFP4 GEMM v25t_1: two-phase exhaustive config search."""
A, _, B_q, B_shuffle, B_scale_sh = data
- return _ext.fwd(A.cuda(), B_q, B_shuffle, B_scale_sh)
+ A = A.contiguous().cuda()
+
+ M, K_elem = A.shape
+ N = B_q.shape[0]
+ K_packed = K_elem // 2
+
+ # Quant
+ quant_result = _ext.do_quant(A)
+ A_fp4 = quant_result[0]
+ A_scale = quant_result[1]
+
+ # Cache B views
+ bsh_dp = B_shuffle.data_ptr()
+ bssh_dp = B_scale_sh.data_ptr()
+ if bsh_dp != _b_cache["bsh_dp"] or bssh_dp != _b_cache["bssh_dp"]:
+ w_raw = B_shuffle.view(torch.uint8) if B_shuffle.dtype != torch.uint8 else B_shuffle
+ _b_cache["w"] = w_raw.reshape(N // 16, K_packed * 16).contiguous()
+ bs_raw = B_scale_sh.view(torch.uint8) if B_scale_sh.dtype != torch.uint8 else B_scale_sh
+ _b_cache["b_scales"] = bs_raw.contiguous()
+ _b_cache["bsh_dp"] = bsh_dp
+ _b_cache["bssh_dp"] = bssh_dp
+
+ w = _b_cache["w"]
+ b_scales = _b_cache["b_scales"]
+
+ # --- Search only for shapes listed in SEARCH_SHAPES; defaults for rest ---
+ shape_key = (M, N, K_elem)
+ if shape_key not in _CONFIG_CACHE:
+ if shape_key in SEARCH_SHAPES:
+ print(f"[v25t_1] exhaustive search for M={M} N={N} K={K_elem} ...", flush=True)
+ try:
+ best = _search_config(M, N, K_elem, A_fp4, A_scale, w, b_scales)
+ _CONFIG_CACHE[shape_key] = best
+ cm = best.get("cache_modifier", None)
+ print(f"[v25t_1] BEST M={M} N={N} K={K_elem}: "
+ f"BM={best['BLOCK_SIZE_M']} BN={best['BLOCK_SIZE_N']} "
+ f"BK={best['BLOCK_SIZE_K']} warps={best['num_warps']} "
+ f"stages={best['num_stages']} wpe={best['waves_per_eu']} "
+ f"GSM={best['GROUP_SIZE_M']} splitK={best['NUM_KSPLIT']} "
+ f"cache={cm}", flush=True)
+ except Exception as e:
+ print(f"[v25t_1] search failed: {e}, using default", flush=True)
+ _CONFIG_CACHE[shape_key] = _DEFAULTS.get(
+ shape_key, _cfg(16, 64, 256, 4, 2, 2, 1, None, 1)
+ )
+ else:
+ # Test shape — use defaults, no search
+ _CONFIG_CACHE[shape_key] = _DEFAULTS.get(
+ shape_key, _cfg(16, 64, 256, 4, 2, 2, 1, None, 1)
+ )
+ print(f"[v25t_1] test shape M={M} N={N} K={K_elem}, using default", flush=True)
+
+ config = _CONFIG_CACHE[shape_key]
+ NUM_KSPLIT = config.get("NUM_KSPLIT", 1)
+
+ _ws.ensure(M, N, NUM_KSPLIT, A.device)
+
+ return _run_gemm(config, A_fp4, A_scale, w, b_scales, M, N, K_elem, K_packed,
+ _ws.y, _ws.y_pp)
scrolls · 1044 diff lines total

Best evidence level for this revision: reported

JSON