Skip to content
KernelIndex
Search⌘K

submission 754280

Hamza · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

hybrid-GEMM.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-754280?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.00µs
#16 of 1143
2026-04-07

Reported · How evidence levels are derived →

Source and license

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

Techniques

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

fp4Fused BF16-to-FP4 quantization + scaled GEMM with pre-shuffled weights.
split-kPer-shape tuned tile/split-K parameters for MI355X (256 CUs).

Kernel source

hybrid-GEMM.py443 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
Fused BF16-to-FP4 quantization + scaled GEMM with pre-shuffled weights.
Hardware-accelerated quantization via v_cvt_scalef32_pk_fp4_bf16.
Per-shape tuned tile/split-K parameters for MI355X (256 CUs).
"""
from task import input_t, output_t
import torch
import triton
import triton.language as tl
import sys as _sys
import gc as _gc

torch.set_grad_enabled(False)
_gc.disable()
_sys.setswitchinterval(1.0)


@triton.jit
def _quant_block_fp4(
    inp_bf16,
    TILE_M: tl.constexpr,
    TILE_K: tl.constexpr,
):
    """Convert BF16 tile to packed MXFP4 using hardware pk_fp4 instruction."""
    GRP_SZ: tl.constexpr = 32
    N_GROUPS: tl.constexpr = TILE_K // GRP_SZ

    vals_f32 = inp_bf16.to(tl.float32).reshape(TILE_M, N_GROUPS, GRP_SZ)

    # Block-wise absolute max with rounding to nearest power-of-2
    peak = tl.max(tl.abs(vals_f32), axis=-1, keep_dims=True)
    peak = peak.to(tl.int32, bitcast=True)
    peak = (peak + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
    log2_peak = ((peak >> 23) & 0xFF).to(tl.int32) - 127
    exp_unbiased = log2_peak - 2
    exp_unbiased = tl.minimum(tl.maximum(exp_unbiased, -127), 127)
    scale_e8m0 = exp_unbiased.to(tl.uint8) + 127

    # Build IEEE754 float divisor: 2^unbiased for hw instruction
    div_bits = (exp_unbiased.to(tl.int32) + 127).to(tl.uint32) << 23
    divisor = div_bits.to(tl.float32, bitcast=True)  # [M, N_GROUPS, 1]

    # Expand divisor to per-pair level
    div_full = tl.broadcast_to(divisor, (TILE_M, N_GROUPS, GRP_SZ))
    div_full = div_full.reshape(TILE_M, TILE_K)
    div_pairs = div_full.reshape(TILE_M, TILE_K // 2, 2)
    div_even, _ = tl.split(div_pairs)
    div_per_pair = div_even.reshape(TILE_M, TILE_K // 2)

    # Interleave adjacent BF16 elements into uint32 for hw conversion
    inp_u16 = inp_bf16.to(tl.uint16, bitcast=True).reshape(TILE_M, TILE_K // 2, 2)
    lo_half, hi_half = tl.split(inp_u16)
    packed_u32 = lo_half.to(tl.uint32) | (hi_half.to(tl.uint32) << 16)
    packed_u32 = packed_u32.reshape(TILE_M, TILE_K // 2)

    # Hardware FP4 pack-convert
    fp4_raw = tl.inline_asm_elementwise(
        "v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2",
        "=v, v, v",
        [packed_u32, div_per_pair],
        dtype=tl.uint32,
        is_pure=True,
        pack=1,
    )
    fp4_bytes = (fp4_raw & 0xFF).to(tl.uint8)
    fp4_bytes = fp4_bytes.reshape(TILE_M, TILE_K // 2)

    return fp4_bytes, scale_e8m0.reshape(TILE_M, N_GROUPS)


@triton.heuristics(
    {
        "ALIGNED_K": lambda args: (args["K"] % (args["TILE_K"] // 2) == 0)
        and (args["SK_TILE"] % args["TILE_K"] == 0)
        and (args["K"] % (args["SK_TILE"] // 2) == 0),
    }
)
@triton.jit
def _matmul_fused_kernel(
    inp_ptr, wt_ptr, out_ptr, wsc_ptr,
    M, N, K,
    stride_im, stride_ik,
    stride_wn, stride_wk,
    stride_ok, stride_om, stride_on,
    stride_sn, stride_sk,
    TILE_M: tl.constexpr,
    TILE_N: tl.constexpr,
    TILE_K: tl.constexpr,
    GROUP_M: tl.constexpr,
    N_SPLITS: tl.constexpr,
    SK_TILE: tl.constexpr,
    ALIGNED_K: tl.constexpr,
    num_warps: tl.constexpr,
    num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr,
    matrix_instr_nonkdim: tl.constexpr,
    load_modifier: tl.constexpr,
):
    tl.assume(stride_im > 0)
    tl.assume(stride_ik > 0)
    tl.assume(stride_wn > 0)
    tl.assume(stride_wk > 0)
    tl.assume(stride_om > 0)
    tl.assume(stride_on > 0)
    tl.assume(stride_sn > 0)
    tl.assume(stride_sk > 0)

    SCALE_GRP: tl.constexpr = 32
    n_tiles_m = tl.cdiv(M, TILE_M)
    n_tiles_n = tl.cdiv(N, TILE_N)

    flat_pid = tl.program_id(axis=0)
    split_id = flat_pid % N_SPLITS
    tile_pid = flat_pid // N_SPLITS

    # Tile assignment: grouped swizzle for single-split, linear for multi-split
    if N_SPLITS == 1:
        tiles_per_grp = GROUP_M * n_tiles_n
        grp = tile_pid // tiles_per_grp
        first_m = grp * GROUP_M
        grp_sz = min(n_tiles_m - first_m, GROUP_M)
        tile_m = first_m + ((tile_pid % tiles_per_grp) % grp_sz)
        tile_n = (tile_pid % tiles_per_grp) // grp_sz
    else:
        tile_m = tile_pid // n_tiles_n
        tile_n = tile_pid % n_tiles_n

    tl.assume(tile_m >= 0)
    tl.assume(tile_n >= 0)
    tl.assume(split_id >= 0)

    if (split_id * SK_TILE // 2) < K:
        k_iters = tl.cdiv(SK_TILE // 2, TILE_K // 2)

        # A: BF16 input [M, 2*K]
        row_a = (tile_m * TILE_M + tl.arange(0, TILE_M)) % M
        col_a = split_id * SK_TILE + tl.arange(0, TILE_K)
        ptrs_a = inp_ptr + (row_a[:, None] * stride_im + col_a[None, :] * stride_ik)

        # B: pre-shuffled FP4 weights [N//16, K_packed*16]
        shuf_range = tl.arange(0, (TILE_K // 2) * 16)
        shuf_base = split_id * (SK_TILE // 2) * 16 + shuf_range
        row_b = (tile_n * (TILE_N // 16) + tl.arange(0, TILE_N // 16)) % (N // 16)
        ptrs_b = wt_ptr + (row_b[:, None] * stride_wn + shuf_base[None, :] * stride_wk)

        # B scales: shuffled E8M0 layout
        row_s = (tile_n * TILE_N + tl.arange(0, TILE_N // 32) * 32)
        col_s = (split_id * (SK_TILE // SCALE_GRP) * 32) + tl.arange(
            0, TILE_K // SCALE_GRP * 32
        )
        ptrs_s = wsc_ptr + row_s[:, None] * stride_sn + col_s[None, :] * stride_sk

        acc = tl.zeros((TILE_M, TILE_N), dtype=tl.float32)

        for ki in range(split_id * k_iters, (split_id + 1) * k_iters):
            # Issue all loads before compute for memory-level parallelism
            if ALIGNED_K:
                a_tile = tl.load(ptrs_a, eviction_policy="evict_last")
                s_raw = tl.load(ptrs_s, cache_modifier=load_modifier)
                b_raw = tl.load(ptrs_b, cache_modifier=load_modifier)
            else:
                k_off = (ki - split_id * k_iters) * TILE_K
                a_tile = tl.load(
                    ptrs_a,
                    mask=tl.arange(0, TILE_K)[None, :] < (2 * K - split_id * SK_TILE - k_off),
                    other=0.0,
                    eviction_policy="evict_last",
                )
                s_raw = tl.load(ptrs_s, cache_modifier=load_modifier)
                b_raw = tl.load(
                    ptrs_b,
                    mask=shuf_range[None, :] < ((K - (split_id * (SK_TILE // 2) + (ki - split_id * k_iters) * (TILE_K // 2))) * 16),
                    other=0,
                    cache_modifier=load_modifier,
                )

            # On-the-fly A quantization
            inp_q, inp_sc = _quant_block_fp4(a_tile, TILE_M, TILE_K)

            # Reconstruct B scale layout from shuffled storage
            b_sc = (
                s_raw
                .reshape(
                    TILE_N // 32,
                    TILE_K // SCALE_GRP // 8,
                    4, 16, 2, 2, 1,
                )
                .permute(0, 5, 3, 1, 4, 2, 6)
                .reshape(TILE_N, TILE_K // SCALE_GRP)
            )

            # Reconstruct B tile from shuffled storage
            b_tile = (
                b_raw.reshape(1, TILE_N // 16, TILE_K // 64, 2, 16, 16)
                .permute(0, 1, 4, 2, 3, 5)
                .reshape(TILE_N, TILE_K // 2)
                .trans(1, 0)
            )

            acc = tl.dot_scaled(
                inp_q, inp_sc, "e2m1", b_tile, b_sc, "e2m1", acc,
                fast_math=True,
            )

            ptrs_a += TILE_K * stride_ik
            ptrs_b += (TILE_K // 2) * 16 * stride_wk
            ptrs_s += TILE_K * stride_sk

        out_vals = acc.to(out_ptr.type.element_ty)

        row_o = tile_m * TILE_M + tl.arange(0, TILE_M).to(tl.int64)
        col_o = tile_n * TILE_N + tl.arange(0, TILE_N).to(tl.int64)
        ptrs_o = (
            out_ptr
            + stride_om * row_o[:, None]
            + stride_on * col_o[None, :]
            + split_id * stride_ok
        )
        mask_o = (row_o[:, None] < M) & (col_o[None, :] < N)
        tl.store(ptrs_o, out_vals, mask=mask_o, cache_modifier=".wt")


@triton.jit
def _partial_sum_kernel(
    partials_ptr, final_ptr, M, N,
    stride_pk, stride_pm, stride_pn,
    stride_fm, stride_fn,
    RED_M: tl.constexpr, RED_N: tl.constexpr,
    TRUE_SPLITS: tl.constexpr, MAX_SPLITS: tl.constexpr,
):
    """Reduce split-K partial results into final output."""
    pid_m = tl.program_id(axis=0)
    pid_n = tl.program_id(axis=1)
    rows = (pid_m * RED_M + tl.arange(0, RED_M)) % M
    cols = (pid_n * RED_N + tl.arange(0, RED_N)) % N

    base = (
        partials_ptr
        + (rows[:, None] * stride_pm)
        + (cols[None, :] * stride_pn)
    )
    total = tl.load(base).to(tl.float32)
    for s in tl.static_range(1, MAX_SPLITS):
        if s < TRUE_SPLITS:
            total += tl.load(base + s * stride_pk).to(tl.float32)

    out_vals = total.to(final_ptr.type.element_ty)
    out_ptrs = (
        final_ptr
        + (rows[:, None] * stride_fm)
        + (cols[None, :] * stride_fn)
    )
    tl.store(out_ptrs, out_vals)


# ---------------------------------------------------------------------------
# Split-K adjustment for divisibility constraints
# ---------------------------------------------------------------------------
def _adjust_splitk(K, BK, n_splits):
    sk_tile = (
        triton.cdiv((2 * triton.cdiv(K, n_splits)), BK) * BK
    )
    while n_splits > 1 and BK > 16:
        if (
            K % (sk_tile // 2) == 0
            and sk_tile % BK == 0
            and K % (BK // 2) == 0
        ):
            break
        elif K % (sk_tile // 2) != 0 and n_splits > 1:
            n_splits = n_splits // 2
        elif sk_tile % BK != 0:
            if n_splits > 1:
                n_splits = n_splits // 2
            elif BK > 16:
                BK = BK // 2
        elif K % (BK // 2) != 0 and BK > 16:
            BK = BK // 2
        else:
            break
        sk_tile = (
            triton.cdiv((2 * triton.cdiv(K, n_splits)), BK) * BK
        )
    n_splits = triton.cdiv(K, (sk_tile // 2))
    return sk_tile, BK, n_splits


# ---------------------------------------------------------------------------
# Per-shape tuned configurations (MI355X, 256 CUs)
# ---------------------------------------------------------------------------
_TUNED_PARAMS = {
    (4, 2880, 512): {"TILE_M": 4, "TILE_N": 128, "TILE_K": 256, "GROUP_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 1, "matrix_instr_nonkdim": 16, "load_modifier": None, "N_SPLITS": 1},
    (16, 2112, 7168): {"TILE_M": 16, "TILE_N": 128, "TILE_K": 512, "GROUP_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "load_modifier": ".cg", "N_SPLITS": 14},
    (32, 4096, 512): {"TILE_M": 16, "TILE_N": 32, "TILE_K": 256, "GROUP_M": 1, "num_warps": 4, "num_stages": 3, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "load_modifier": ".cg", "N_SPLITS": 1},
    (32, 2880, 512): {"TILE_M": 8, "TILE_N": 128, "TILE_K": 256, "GROUP_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "load_modifier": None, "N_SPLITS": 1},
    (64, 7168, 2048): {"TILE_M": 16, "TILE_N": 128, "TILE_K": 256, "GROUP_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "load_modifier": ".cg", "N_SPLITS": 1},
    (256, 3072, 1536): {"TILE_M": 16, "TILE_N": 256, "TILE_K": 512, "GROUP_M": 1, "num_warps": 8, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "load_modifier": None, "N_SPLITS": 1},
}

_FALLBACK_PARAMS = {
    "TILE_M": 16, "TILE_N": 32, "TILE_K": 256, "GROUP_M": 1,
    "num_warps": 2, "num_stages": 2, "waves_per_eu": 0,
    "matrix_instr_nonkdim": 16, "load_modifier": ".cg", "N_SPLITS": 1,
}


# ---------------------------------------------------------------------------
# Caches for buffers, configs, and launch grids
# ---------------------------------------------------------------------------
_out_pool = {}
_param_pool = {}
_grid_pool = {}
_operand_pool = {}


def _get_output_bufs(m, n, n_splits, dev):
    key = (m, n, n_splits)
    if key not in _out_pool:
        final = torch.empty((m, n), dtype=torch.bfloat16, device=dev)
        partials = (
            torch.empty((n_splits, m, n), dtype=torch.float32, device=dev)
            if n_splits > 1 else None
        )
        _out_pool[key] = (final, partials)
    return _out_pool[key]


def _resolve_params(m, n, k):
    key = (m, n, k)
    if key not in _param_pool:
        params = _TUNED_PARAMS.get(key, _FALLBACK_PARAMS).copy()
        k_half = k // 2
        if params["N_SPLITS"] > 1:
            sk_tile, tile_k, n_splits = _adjust_splitk(
                k_half, params["TILE_K"], params["N_SPLITS"]
            )
            params["SK_TILE"] = sk_tile
            params["TILE_K"] = tile_k
            params["N_SPLITS"] = n_splits
        else:
            params["SK_TILE"] = 2 * k_half
            params["N_SPLITS"] = 1
        if params["TILE_K"] >= 2 * k_half:
            params["TILE_K"] = triton.next_power_of_2(2 * k_half)
            params["SK_TILE"] = 2 * k_half
            params["N_SPLITS"] = 1
        params["TILE_N"] = max(params["TILE_N"], 32)
        _param_pool[key] = params
    return _param_pool[key]


def _reshape_operands(B_data, B_scales, n, k_half):
    """Reshape pre-shuffled weight and scale tensors for kernel indexing."""
    key = B_data.data_ptr()
    if key not in _operand_pool:
        b_flat = B_data.view(torch.uint8).reshape(n // 16, k_half * 16)
        s_flat = B_scales.view(torch.uint8)
        _operand_pool[key] = (b_flat, s_flat)
    return _operand_pool[key]


def _build_grid(m, n, k, dev):
    """Precompute all launch parameters for a given problem shape."""
    key = (m, n, k)
    if key not in _grid_pool:
        params = _resolve_params(m, n, k)
        k_half = k // 2
        ns = params["N_SPLITS"]

        final, partials = _get_output_bufs(m, n, ns, dev)

        grid = (
            ns
            * triton.cdiv(m, params["TILE_M"])
            * triton.cdiv(n, params["TILE_N"]),
        )

        if ns == 1:
            sk_o, sm_o, sn_o = 0, final.stride(0), final.stride(1)
        else:
            sk_o = partials.stride(0)
            sm_o = partials.stride(1)
            sn_o = partials.stride(2)

        info = {
            'params': params,
            'k_half': k_half,
            'grid': grid,
            'ns': ns,
            'sk_o': sk_o,
            'sm_o': sm_o,
            'sn_o': sn_o,
        }

        if ns > 1:
            info['red_grid'] = (triton.cdiv(m, 16), triton.cdiv(n, 64))
            info['true_splits'] = triton.cdiv(k_half, (params["SK_TILE"] // 2))
            info['max_splits'] = triton.next_power_of_2(ns)

        _grid_pool[key] = info
    return _grid_pool[key]


def _execute_matmul(inp_mat, wt_data, wt_scales, m, n, k):
    """Run fused quantize + GEMM with optional split-K reduction."""
    info = _build_grid(m, n, k, inp_mat.device)
    final, partials = _get_output_bufs(m, n, info['ns'], inp_mat.device)

    b_flat, s_flat = _reshape_operands(wt_data, wt_scales, n, info['k_half'])

    _matmul_fused_kernel[info['grid']](
        inp_mat, b_flat,
        final if info['ns'] == 1 else partials,
        s_flat,
        m, n, info['k_half'],
        inp_mat.stride(0), inp_mat.stride(1),
        b_flat.stride(0), b_flat.stride(1),
        info['sk_o'], info['sm_o'], info['sn_o'],
        s_flat.stride(0), s_flat.stride(1),
        **info['params'],
    )

    if info['ns'] > 1:
        _partial_sum_kernel[info['red_grid']](
            partials, final, m, n,
            partials.stride(0), partials.stride(1), partials.stride(2),
            final.stride(0), final.stride(1),
            16, 64,
            info['true_splits'], info['max_splits'],
        )

    return final


def custom_kernel(data: input_t) -> output_t:
    A = data[0]
    return _execute_matmul(
        A, data[3], data[4], A.shape[0], data[1].shape[0], A.shape[1]
    )
scrolls · 443 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 734513.

#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
- # v64_wavesched — aiter update + eviction_policy + TRITON_HIP_ENABLE_WAVE_SCHEDULING=1
- # Best of session 55: marginal but consistent improvement over v62
+ """
+ Fused BF16-to-FP4 quantization + scaled GEMM with pre-shuffled weights.
+ Hardware-accelerated quantization via v_cvt_scalef32_pk_fp4_bf16.
+ Per-shape tuned tile/split-K parameters for MI355X (256 CUs).
+ """
+ from task import input_t, output_t
+ import torch
+ import triton
+ import triton.language as tl
+ import sys as _sys
+ import gc as _gc
- import os as _os
- import sys as _isys
- import subprocess as _isp
- import time as _itime
- _os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
- _os.environ.setdefault("CXX", "clang++")
+ torch.set_grad_enabled(False)
+ _gc.disable()
+ _sys.setswitchinterval(1.0)
- _IT0 = _itime.time()
- _pe = lambda msg: print(msg, file=_isys.stderr, flush=True)
- # ============================================================
- # PHASE 0: Update aiter to origin/main (has MI355X tuned configs)
- # ============================================================
- _AITER_DIR = '/home/runner/aiter'
- _AITER_UPDATED = False
- try:
- _pe("[v62] Fetching origin/main...")
- _r = _isp.run(['git', '-C', _AITER_DIR, 'fetch', 'origin', 'main'],
- capture_output=True, text=True, timeout=60)
- _pe(f" fetch: rc={_r.returncode}")
+ @triton.jit
+ def _quant_block_fp4(
+ inp_bf16,
+ TILE_M: tl.constexpr,
+ TILE_K: tl.constexpr,
+ ):
+ """Convert BF16 tile to packed MXFP4 using hardware pk_fp4 instruction."""
+ GRP_SZ: tl.constexpr = 32
+ N_GROUPS: tl.constexpr = TILE_K // GRP_SZ
- # Save current HEAD for rollback
- _r0 = _isp.run(['git', '-C', _AITER_DIR, 'rev-parse', 'HEAD'],
- capture_output=True, text=True, timeout=5)
- _OLD_HEAD = _r0.stdout.strip()
- _pe(f" old HEAD: {_OLD_HEAD[:12]}")
+ vals_f32 = inp_bf16.to(tl.float32).reshape(TILE_M, N_GROUPS, GRP_SZ)
- # Checkout origin/main
- _r = _isp.run(['git', '-C', _AITER_DIR, 'checkout', 'origin/main'],
- capture_output=True, text=True, timeout=30)
- _pe(f" checkout origin/main: rc={_r.returncode}")
- if _r.stderr.strip():
- _pe(f" checkout err: {_r.stderr.strip()[:200]}")
+ # Block-wise absolute max with rounding to nearest power-of-2
+ peak = tl.max(tl.abs(vals_f32), axis=-1, keep_dims=True)
+ peak = peak.to(tl.int32, bitcast=True)
+ peak = (peak + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
+ log2_peak = ((peak >> 23) & 0xFF).to(tl.int32) - 127
+ exp_unbiased = log2_peak - 2
+ exp_unbiased = tl.minimum(tl.maximum(exp_unbiased, -127), 127)
+ scale_e8m0 = exp_unbiased.to(tl.uint8) + 127
- if _r.returncode == 0:
- _r2 = _isp.run(['git', '-C', _AITER_DIR, 'log', '--oneline', '-5'],
- capture_output=True, text=True, timeout=5)
- _pe(f" new HEAD:\n{_r2.stdout.strip()}")
- _AITER_UPDATED = True
- else:
- _pe(" checkout FAILED, staying on old HEAD")
- except Exception as _e:
- _pe(f" [aiter update] FAILED: {_e}")
+ # Build IEEE754 float divisor: 2^unbiased for hw instruction
+ div_bits = (exp_unbiased.to(tl.int32) + 127).to(tl.uint32) << 23
+ divisor = div_bits.to(tl.float32, bitcast=True) # [M, N_GROUPS, 1]
- # PHASE 0b removed — eviction_policy now applied via in-memory _unsafe_update_src (Patch 4)
- _KERN_PATCHED = False
+ # Expand divisor to per-pair level
+ div_full = tl.broadcast_to(divisor, (TILE_M, N_GROUPS, GRP_SZ))
+ div_full = div_full.reshape(TILE_M, TILE_K)
+ div_pairs = div_full.reshape(TILE_M, TILE_K // 2, 2)
+ div_even, _ = tl.split(div_pairs)
+ div_per_pair = div_even.reshape(TILE_M, TILE_K // 2)
- # ============================================================
- # PHASE 0c: Read new tuned configs if available
- # ============================================================
- try:
- _cfg_path = '/home/runner/aiter/aiter/configs/a4w4_blockscale_tuned_gemm.csv'
- if _os.path.exists(_cfg_path):
- with open(_cfg_path) as _f:
- _cfg_lines = _f.readlines()
- _pe(f"[v62] Tuned config: {len(_cfg_lines)} lines")
- # Print first few + last few lines
- for _l in _cfg_lines[:3]:
- _pe(f" {_l.rstrip()}")
- if len(_cfg_lines) > 6:
- _pe(" ...")
- for _l in _cfg_lines[-3:]:
- _pe(f" {_l.rstrip()}")
+ # Interleave adjacent BF16 elements into uint32 for hw conversion
+ inp_u16 = inp_bf16.to(tl.uint16, bitcast=True).reshape(TILE_M, TILE_K // 2, 2)
+ lo_half, hi_half = tl.split(inp_u16)
+ packed_u32 = lo_half.to(tl.uint32) | (hi_half.to(tl.uint32) << 16)
+ packed_u32 = packed_u32.reshape(TILE_M, TILE_K // 2)
- # Check for MI355X-specific or new entries
- _mi355_lines = [l for l in _cfg_lines if '256' in l.split(',')[0:1]]
- _pe(f" entries with 256 CUs: {len(_mi355_lines)}")
- except Exception as _e:
- _pe(f" [tuned cfg] {_e}")
+ # Hardware FP4 pack-convert
+ fp4_raw = tl.inline_asm_elementwise(
+ "v_cvt_scalef32_pk_fp4_bf16 $0, $1, $2",
+ "=v, v, v",
+ [packed_u32, div_per_pair],
+ dtype=tl.uint32,
+ is_pure=True,
+ pack=1,
+ )
+ fp4_bytes = (fp4_raw & 0xFF).to(tl.uint8)
+ fp4_bytes = fp4_bytes.reshape(TILE_M, TILE_K // 2)
- _pe(f"[v62] Init phase: {_itime.time()-_IT0:.1f}s, updated={_AITER_UPDATED}, patched={_KERN_PATCHED}")
- del _isp, _itime, _pe, _IT0
+ return fp4_bytes, scale_e8m0.reshape(TILE_M, N_GROUPS)
- import uuid as _uuid
- _os.environ["TRITON_CACHE_DIR"] = f"/tmp/_triton_v64_{_uuid.uuid4().hex[:8]}"
- _os.environ["TRITON_HIP_ENABLE_WAVE_SCHEDULING"] = "1"
- _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))
- # Always use ONLY our CSV — prevents module_gemm_common/a4w4_asm builds (5s+ overhead)
- # Our Triton preshuffle kernel bypasses the CSV entirely for actual computation
- _os.environ["AITER_CONFIG_GEMM_A4W4"] = _CSV_PATH
+ @triton.heuristics(
+ {
+ "ALIGNED_K": lambda args: (args["K"] % (args["TILE_K"] // 2) == 0)
+ and (args["SK_TILE"] % args["TILE_K"] == 0)
+ and (args["K"] % (args["SK_TILE"] // 2) == 0),
+ }
+ )
+ @triton.jit
+ def _matmul_fused_kernel(
+ inp_ptr, wt_ptr, out_ptr, wsc_ptr,
+ M, N, K,
+ stride_im, stride_ik,
+ stride_wn, stride_wk,
+ stride_ok, stride_om, stride_on,
+ stride_sn, stride_sk,
+ TILE_M: tl.constexpr,
+ TILE_N: tl.constexpr,
+ TILE_K: tl.constexpr,
+ GROUP_M: tl.constexpr,
+ N_SPLITS: tl.constexpr,
+ SK_TILE: tl.constexpr,
+ ALIGNED_K: tl.constexpr,
+ num_warps: tl.constexpr,
+ num_stages: tl.constexpr,
+ waves_per_eu: tl.constexpr,
+ matrix_instr_nonkdim: tl.constexpr,
+ load_modifier: tl.constexpr,
+ ):
+ tl.assume(stride_im > 0)
+ tl.assume(stride_ik > 0)
+ tl.assume(stride_wn > 0)
+ tl.assume(stride_wk > 0)
+ tl.assume(stride_om > 0)
+ tl.assume(stride_on > 0)
+ tl.assume(stride_sn > 0)
+ tl.assume(stride_sk > 0)
- import torch
- torch.set_grad_enabled(False)
- import triton
- import triton.language as tl
- import sys as _sys
- import time as _time
- import gc as _gc
- _sys.setswitchinterval(1.0)
+ SCALE_GRP: tl.constexpr = 32
+ n_tiles_m = tl.cdiv(M, TILE_M)
+ n_tiles_n = tl.cdiv(N, TILE_N)
- # Import with rollback safety — if updated aiter breaks, revert to old HEAD
- try:
- 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,
- )
- print("[v62] aiter import OK", file=_sys.stderr, flush=True)
- except Exception as _import_err:
- print(f"[v62] aiter import FAILED: {_import_err}, rolling back...", file=_sys.stderr, flush=True)
- import subprocess as _rbsp
- try:
- _rbsp.run(['git', '-C', '/home/runner/aiter', 'checkout', _OLD_HEAD],
- capture_output=True, text=True, timeout=15)
- import importlib
- # Re-import with old code
- 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,
- )
- print("[v62] rollback OK, using old aiter", file=_sys.stderr, flush=True)
- _AITER_UPDATED = False
- except Exception as _rb_err:
- print(f"[v62] rollback FAILED: {_rb_err}", file=_sys.stderr, flush=True)
- raise _import_err
- del _rbsp
+ flat_pid = tl.program_id(axis=0)
+ split_id = flat_pid % N_SPLITS
+ tile_pid = flat_pid // N_SPLITS
- from task import input_t, output_t
+ # Tile assignment: grouped swizzle for single-split, linear for multi-split
+ if N_SPLITS == 1:
+ tiles_per_grp = GROUP_M * n_tiles_n
+ grp = tile_pid // tiles_per_grp
+ first_m = grp * GROUP_M
+ grp_sz = min(n_tiles_m - first_m, GROUP_M)
+ tile_m = first_m + ((tile_pid % tiles_per_grp) % grp_sz)
+ tile_n = (tile_pid % tiles_per_grp) // grp_sz
+ else:
+ tile_m = tile_pid // n_tiles_n
+ tile_n = tile_pid % n_tiles_n
- # --- Monkey-patch heuristics ---
- try:
- # v55: restore default GRID_MN (tile grouping for L2 locality)
- _gemm_a16wfp4_preshuffle_kernel.values['EVEN_K'] = lambda args: True
- print("[patch] EVEN_K → True (GRID_MN = default)", file=_sys.stderr, flush=True)
- except Exception as _e:
- print(f"[patch] heuristics failed: {_e}", file=_sys.stderr, flush=True)
+ tl.assume(tile_m >= 0)
+ tl.assume(tile_n >= 0)
+ tl.assume(split_id >= 0)
- _os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
+ if (split_id * SK_TILE // 2) < K:
+ k_iters = tl.cdiv(SK_TILE // 2, TILE_K // 2)
- # --- 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
+ # A: BF16 input [M, 2*K]
+ row_a = (tile_m * TILE_M + tl.arange(0, TILE_M)) % M
+ col_a = split_id * SK_TILE + tl.arange(0, TILE_K)
+ ptrs_a = inp_ptr + (row_a[:, None] * stride_im + col_a[None, :] * stride_ik)
- # 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
+ # B: pre-shuffled FP4 weights [N//16, K_packed*16]
+ shuf_range = tl.arange(0, (TILE_K // 2) * 16)
+ shuf_base = split_id * (SK_TILE // 2) * 16 + shuf_range
+ row_b = (tile_n * (TILE_N // 16) + tl.arange(0, TILE_N // 16)) % (N // 16)
+ ptrs_b = wt_ptr + (row_b[:, None] * stride_wn + shuf_base[None, :] * stride_wk)
- x = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
+ # B scales: shuffled E8M0 layout
+ row_s = (tile_n * TILE_N + tl.arange(0, TILE_N // 32) * 32)
+ col_s = (split_id * (SK_TILE // SCALE_GRP) * 32) + tl.arange(
+ 0, TILE_K // SCALE_GRP * 32
+ )
+ ptrs_s = wsc_ptr + row_s[:, None] * stride_sn + col_s[None, :] * stride_sk
- # 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
+ acc = tl.zeros((TILE_M, TILE_N), dtype=tl.float32)
- # 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)
+ for ki in range(split_id * k_iters, (split_id + 1) * k_iters):
+ # Issue all loads before compute for memory-level parallelism
+ if ALIGNED_K:
+ a_tile = tl.load(ptrs_a, eviction_policy="evict_last")
+ s_raw = tl.load(ptrs_s, cache_modifier=load_modifier)
+ b_raw = tl.load(ptrs_b, cache_modifier=load_modifier)
+ else:
+ k_off = (ki - split_id * k_iters) * TILE_K
+ a_tile = tl.load(
+ ptrs_a,
+ mask=tl.arange(0, TILE_K)[None, :] < (2 * K - split_id * SK_TILE - k_off),
+ other=0.0,
+ eviction_policy="evict_last",
+ )
+ s_raw = tl.load(ptrs_s, cache_modifier=load_modifier)
+ b_raw = tl.load(
+ ptrs_b,
+ mask=shuf_range[None, :] < ((K - (split_id * (SK_TILE // 2) + (ki - split_id * k_iters) * (TILE_K // 2))) * 16),
+ other=0,
+ cache_modifier=load_modifier,
+ )
- # E8M0 scale bytes for output
- bs_e8m0 = (scale_e8m0_unbiased + 127).to(tl.float32).to(tl.uint8)
+ # On-the-fly A quantization
+ inp_q, inp_sc = _quant_block_fp4(a_tile, TILE_M, TILE_K)
- # 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)
+ # Reconstruct B scale layout from shuffled storage
+ b_sc = (
+ s_raw
+ .reshape(
+ TILE_N // 32,
+ TILE_K // SCALE_GRP // 8,
+ 4, 16, 2, 2, 1,
+ )
+ .permute(0, 5, 3, 1, 4, 2, 6)
+ .reshape(TILE_N, TILE_K // SCALE_GRP)
+ )
- # 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]
+ # Reconstruct B tile from shuffled storage
+ b_tile = (
+ b_raw.reshape(1, TILE_N // 16, TILE_K // 64, 2, 16, 16)
+ .permute(0, 1, 4, 2, 3, 5)
+ .reshape(TILE_N, TILE_K // 2)
+ .trans(1, 0)
+ )
- # 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,
- )
+ acc = tl.dot_scaled(
+ inp_q, inp_sc, "e2m1", b_tile, b_sc, "e2m1", acc,
+ fast_math=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)
+ ptrs_a += TILE_K * stride_ik
+ ptrs_b += (TILE_K // 2) * 16 * stride_wk
+ ptrs_s += TILE_K * stride_sk
- return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
- '''
+ out_vals = acc.to(out_ptr.type.element_ty)
- 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
- # Patch 1: acc=accumulator (avoids extra zero-init)
- _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)'
- )
- # Patch 2: fast_math=True (relaxed FP precision for MFMA scheduling)
- _new_ksrc = _new_ksrc.replace(
- 'acc=accumulator)',
- 'acc=accumulator, fast_math=True)'
- )
- # Patch 3: .wt store modifier (write-through — avoids L2 pollution from output writes)
- _new_ksrc = _new_ksrc.replace(
- 'tl.store(c_ptrs, c, mask=c_mask)',
- 'tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")'
- )
- # Patch 4: eviction_policy for A loads (keep A in L2 for N-tile reuse)
- _evict_count = 0
- if 'a_bf16 = tl.load(a_ptrs)' in _new_ksrc:
- _new_ksrc = _new_ksrc.replace(
- 'a_bf16 = tl.load(a_ptrs)',
- 'a_bf16 = tl.load(a_ptrs, eviction_policy="evict_last")'
+ row_o = tile_m * TILE_M + tl.arange(0, TILE_M).to(tl.int64)
+ col_o = tile_n * TILE_N + tl.arange(0, TILE_N).to(tl.int64)
+ ptrs_o = (
+ out_ptr
+ + stride_om * row_o[:, None]
+ + stride_on * col_o[None, :]
+ + split_id * stride_ok
)
- _evict_count += 1
- # Also patch masked A load (non-EVEN_K path)
- if 'a_bf16 = tl.load(a_ptrs,' in _new_ksrc and 'eviction_policy' not in _new_ksrc.split('a_bf16 = tl.load(a_ptrs,')[1].split(')')[0]:
- # More robust: find "a_bf16 = tl.load(\n a_ptrs,\n mask="
- # and insert eviction_policy before mask
- import re as _re
- _pat = r'(a_bf16 = tl\.load\(\s*\n\s*a_ptrs,)\s*\n(\s*mask=)'
- _rep = r'\1 eviction_policy="evict_last",\n\2'
- _new_ksrc2 = _re.sub(_pat, _rep, _new_ksrc)
- if _new_ksrc2 != _new_ksrc:
- _new_ksrc = _new_ksrc2
- _evict_count += 1
- print(f"[hwfp4] eviction_policy patches: {_evict_count}", file=_sys.stderr, flush=True)
- _n_patches = sum([
- _new_ksrc != _old_ksrc,
- 'fast_math=True' in _new_ksrc,
- 'cache_modifier=".wt"' in _new_ksrc,
- _evict_count > 0,
- ])
- if _new_ksrc != _old_ksrc:
- _jit_fn._unsafe_update_src(_new_ksrc)
- print(f"[hwfp4] Applied hardware quant + {_n_patches} kernel patches", file=_sys.stderr, flush=True)
- else:
- print("[hwfp4] Applied hardware quant, kernel mod FAILED", file=_sys.stderr, flush=True)
+ mask_o = (row_o[:, None] < M) & (col_o[None, :] < N)
+ tl.store(ptrs_o, out_vals, mask=mask_o, cache_modifier=".wt")
- # 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>
+ @triton.jit
+ def _partial_sum_kernel(
+ partials_ptr, final_ptr, M, N,
+ stride_pk, stride_pm, stride_pn,
+ stride_fm, stride_fn,
+ RED_M: tl.constexpr, RED_N: tl.constexpr,
+ TRUE_SPLITS: tl.constexpr, MAX_SPLITS: tl.constexpr,
+ ):
+ """Reduce split-K partial results into final output."""
+ pid_m = tl.program_id(axis=0)
+ pid_n = tl.program_id(axis=1)
+ rows = (pid_m * RED_M + tl.arange(0, RED_M)) % M
+ cols = (pid_n * RED_N + tl.arange(0, RED_N)) % N
- __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);
- }
+ base = (
+ partials_ptr
+ + (rows[:, None] * stride_pm)
+ + (cols[None, :] * stride_pn)
+ )
+ total = tl.load(base).to(tl.float32)
+ for s in tl.static_range(1, MAX_SPLITS):
+ if s < TRUE_SPLITS:
+ total += tl.load(base + s * stride_pk).to(tl.float32)
- 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);
- }
- }
- }
+ out_vals = total.to(final_ptr.type.element_ty)
+ out_ptrs = (
+ final_ptr
+ + (rows[:, None] * stride_fm)
+ + (cols[None, :] * stride_fn)
+ )
+ tl.store(out_ptrs, out_vals)
- __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"],
+ # ---------------------------------------------------------------------------
+ # Split-K adjustment for divisibility constraints
+ # ---------------------------------------------------------------------------
+ def _adjust_splitk(K, BK, n_splits):
+ sk_tile = (
+ triton.cdiv((2 * triton.cdiv(K, n_splits)), BK) * BK
)
- _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):
+ while n_splits > 1 and BK > 16:
+ if (
+ K % (sk_tile // 2) == 0
+ and sk_tile % BK == 0
+ and K % (BK // 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
+ elif K % (sk_tile // 2) != 0 and n_splits > 1:
+ n_splits = n_splits // 2
+ elif sk_tile % BK != 0:
+ if n_splits > 1:
+ n_splits = n_splits // 2
+ elif BK > 16:
+ BK = BK // 2
+ elif K % (BK // 2) != 0 and BK > 16:
+ BK = BK // 2
else:
break
- SPLITK_BLOCK_SIZE = (
- triton.cdiv((2 * triton.cdiv(K, NUM_KSPLIT)), BLOCK_SIZE_K) * BLOCK_SIZE_K
+ sk_tile = (
+ triton.cdiv((2 * triton.cdiv(K, n_splits)), BK) * BK
)
- return SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT
+ n_splits = triton.cdiv(K, (sk_tile // 2))
+ return sk_tile, BK, n_splits
- _CFG_CACHE = {}
+ # ---------------------------------------------------------------------------
+ # Per-shape tuned configurations (MI355X, 256 CUs)
+ # ---------------------------------------------------------------------------
+ _TUNED_PARAMS = {
+ (4, 2880, 512): {"TILE_M": 4, "TILE_N": 128, "TILE_K": 256, "GROUP_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 1, "matrix_instr_nonkdim": 16, "load_modifier": None, "N_SPLITS": 1},
+ (16, 2112, 7168): {"TILE_M": 16, "TILE_N": 128, "TILE_K": 512, "GROUP_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "load_modifier": ".cg", "N_SPLITS": 14},
+ (32, 4096, 512): {"TILE_M": 16, "TILE_N": 32, "TILE_K": 256, "GROUP_M": 1, "num_warps": 4, "num_stages": 3, "waves_per_eu": 3, "matrix_instr_nonkdim": 16, "load_modifier": ".cg", "N_SPLITS": 1},
+ (32, 2880, 512): {"TILE_M": 8, "TILE_N": 128, "TILE_K": 256, "GROUP_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "load_modifier": None, "N_SPLITS": 1},
+ (64, 7168, 2048): {"TILE_M": 16, "TILE_N": 128, "TILE_K": 256, "GROUP_M": 1, "num_warps": 4, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "load_modifier": ".cg", "N_SPLITS": 1},
+ (256, 3072, 1536): {"TILE_M": 16, "TILE_N": 256, "TILE_K": 512, "GROUP_M": 1, "num_warps": 8, "num_stages": 2, "waves_per_eu": 2, "matrix_instr_nonkdim": 16, "load_modifier": None, "N_SPLITS": 1},
+ }
- 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,
- }
+ _FALLBACK_PARAMS = {
+ "TILE_M": 16, "TILE_N": 32, "TILE_K": 256, "GROUP_M": 1,
+ "num_warps": 2, "num_stages": 2, "waves_per_eu": 0,
+ "matrix_instr_nonkdim": 16, "load_modifier": ".cg", "N_SPLITS": 1,
+ }
- 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)
+ # ---------------------------------------------------------------------------
+ # Caches for buffers, configs, and launch grids
+ # ---------------------------------------------------------------------------
+ _out_pool = {}
+ _param_pool = {}
+ _grid_pool = {}
+ _operand_pool = {}
- 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"])
+ def _get_output_bufs(m, n, n_splits, dev):
+ key = (m, n, n_splits)
+ if key not in _out_pool:
+ final = torch.empty((m, n), dtype=torch.bfloat16, device=dev)
+ partials = (
+ torch.empty((n_splits, m, n), dtype=torch.float32, device=dev)
+ if n_splits > 1 else None
+ )
+ _out_pool[key] = (final, partials)
+ return _out_pool[key]
- 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)
- elif _mw > 32:
- _NO_LSR.setdefault(_ck, True) # v53: M>32 also without disable-lsr
+ def _resolve_params(m, n, k):
+ key = (m, n, k)
+ if key not in _param_pool:
+ params = _TUNED_PARAMS.get(key, _FALLBACK_PARAMS).copy()
+ k_half = k // 2
+ if params["N_SPLITS"] > 1:
+ sk_tile, tile_k, n_splits = _adjust_splitk(
+ k_half, params["TILE_K"], params["N_SPLITS"]
+ )
+ params["SK_TILE"] = sk_tile
+ params["TILE_K"] = tile_k
+ params["N_SPLITS"] = n_splits
else:
- _LSR.setdefault(_ck, True)
- if _aw is not None:
- _REDUCE.add((_aw, _nkw))
+ params["SK_TILE"] = 2 * k_half
+ params["N_SPLITS"] = 1
+ if params["TILE_K"] >= 2 * k_half:
+ params["TILE_K"] = triton.next_power_of_2(2 * k_half)
+ params["SK_TILE"] = 2 * k_half
+ params["N_SPLITS"] = 1
+ params["TILE_N"] = max(params["TILE_N"], 32)
+ _param_pool[key] = params
+ return _param_pool[key]
- 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)
+ def _reshape_operands(B_data, B_scales, n, k_half):
+ """Reshape pre-shuffled weight and scale tensors for kernel indexing."""
+ key = B_data.data_ptr()
+ if key not in _operand_pool:
+ b_flat = B_data.view(torch.uint8).reshape(n // 16, k_half * 16)
+ s_flat = B_scales.view(torch.uint8)
+ _operand_pool[key] = (b_flat, s_flat)
+ return _operand_pool[key]
- _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)
+ def _build_grid(m, n, k, dev):
+ """Precompute all launch parameters for a given problem shape."""
+ key = (m, n, k)
+ if key not in _grid_pool:
+ params = _resolve_params(m, n, k)
+ k_half = k // 2
+ ns = params["N_SPLITS"]
- 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)
+ final, partials = _get_output_bufs(m, n, ns, dev)
- _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)
+ grid = (
+ ns
+ * triton.cdiv(m, params["TILE_M"])
+ * triton.cdiv(n, params["TILE_N"]),
+ )
- _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)
+ if ns == 1:
+ sk_o, sm_o, sn_o = 0, final.stride(0), final.stride(1)
+ else:
+ sk_o = partials.stride(0)
+ sm_o = partials.stride(1)
+ sn_o = partials.stride(2)
- 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
+ info = {
+ 'params': params,
+ 'k_half': k_half,
+ 'grid': grid,
+ 'ns': ns,
+ 'sk_o': sk_o,
+ 'sm_o': sm_o,
+ 'sn_o': sn_o,
+ }
- 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)
+ if ns > 1:
+ info['red_grid'] = (triton.cdiv(m, 16), triton.cdiv(n, 64))
+ info['true_splits'] = triton.cdiv(k_half, (params["SK_TILE"] // 2))
+ info['max_splits'] = triton.next_power_of_2(ns)
- _gc.disable()
+ _grid_pool[key] = info
+ return _grid_pool[key]
- # --- 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 _execute_matmul(inp_mat, wt_data, wt_scales, m, n, k):
+ """Run fused quantize + GEMM with optional split-K reduction."""
+ info = _build_grid(m, n, k, inp_mat.device)
+ final, partials = _get_output_bufs(m, n, info['ns'], inp_mat.device)
- 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
+ b_flat, s_flat = _reshape_operands(wt_data, wt_scales, n, info['k_half'])
- cfg, actual_ksplit, nk_pow2, grid_main, grid_reduce, K, BM, BN, BK, KS, SPK, WPE = _get_cfg(M, N, K_real)
+ _matmul_fused_kernel[info['grid']](
+ inp_mat, b_flat,
+ final if info['ns'] == 1 else partials,
+ s_flat,
+ m, n, info['k_half'],
+ inp_mat.stride(0), inp_mat.stride(1),
+ b_flat.stride(0), b_flat.stride(1),
+ info['sk_o'], info['sm_o'], info['sn_o'],
+ s_flat.stride(0), s_flat.stride(1),
+ **info['params'],
+ )
- _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)
+ if info['ns'] > 1:
+ _partial_sum_kernel[info['red_grid']](
+ partials, final, m, n,
+ partials.stride(0), partials.stride(1), partials.stride(2),
+ final.stride(0), final.stride(1),
+ 16, 64,
+ info['true_splits'], info['max_splits'],
+ )
- 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]
+ return final
- 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,
+ def custom_kernel(data: input_t) -> output_t:
+ A = data[0]
+ return _execute_matmul(
+ A, data[3], data[4], A.shape[0], data[1].shape[0], A.shape[1]
)
-
- 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 · 1067 diff lines total

Best evidence level for this revision: reported

JSON