Skip to content
KernelIndex
Search⌘K

submission 545998

johnny.t.shi · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v142_inline_all.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-545998?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
10.3µs
#254 of 1143
2026-03-13

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:2096149077c393eebe86920c689458a72191c9c1b5ecb95e182392c36e4cd53d
license declaredunknown
license concludedunknown
authorsjohnny.t.shi
imported2026-08-15

Techniques

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

num-warps = 8num_warps=8,
split-k_get_splitk_fn = _gemm_mod.get_splitk
stages = 2num_stages=2,
tile-k = 512BLOCK_SIZE_K = 512
tile-m = 16BLOCK_SIZE_M = 16
tile-n = 64BLOCK_SIZE_N = 64 if M <= 16 else 128

Kernel source

submission_v142_inline_all.py493 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
v142: Extreme Python inline optimization.
- Function pointer swap: warmup function swaps to hot path (no if-check per call)
- List-indexed cache by M (no dict lookup in hot path)
- Pre-bound kernel function references (no global lookup)
- @torch.no_grad() to skip autograd overhead
- Minimal tuple unpacking in hot path
- Pre-computed B_q_T strides cached per shape
- All constants pre-computed during warmup/first-see
- ASM path for M>64 with same optimizations
"""
from task import input_t, output_t

import torch
import triton
import triton.language as tl
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import _mxfp4_quant_op
from aiter.ops.gemm_op_a4w4 import get_GEMM_config
from aiter.ops.gemm_op_common import get_padded_m

import aiter.ops.triton.gemm_afp4wfp4 as _gemm_mod
_reduce_kernel = _gemm_mod._gemm_afp4wfp4_reduce_kernel
_get_splitk_fn = _gemm_mod.get_splitk

# Pre-bind dtype constants
_fp4x2 = dtypes.fp4x2
_fp8_e8m0 = dtypes.fp8_e8m0
_bf16 = dtypes.bf16
_uint8 = torch.uint8
_bfloat16 = torch.bfloat16
_float32 = torch.float32

# Pre-bind torch functions to avoid global lookups
_torch_empty = torch.empty
_torch_full = torch.full
_triton_cdiv = triton.cdiv
_triton_np2 = triton.next_power_of_2


@triton.jit
def _remap_xcd(pid, num_pids, NUM_XCDS: tl.constexpr):
    chunk_size = tl.cdiv(num_pids, NUM_XCDS)
    xcd = pid % NUM_XCDS
    pid_in_xcd = pid // NUM_XCDS
    return xcd * chunk_size + pid_in_xcd


@triton.jit
def _pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M: tl.constexpr):
    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 = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
    pid_m = first_pid_m + (pid % num_pid_in_group) % group_size_m
    pid_n = (pid % num_pid_in_group) // group_size_m
    return pid_m, pid_n


@triton.jit
def _fused_quant_gemm_kernel(
    a_ptr, b_ptr, c_ptr, b_scales_ptr,
    M, N, K_real,
    stride_am, stride_ak,
    stride_bk, stride_bn,
    stride_ck, stride_cm, stride_cn,
    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,
    QUANT_BLOCK: tl.constexpr,
):
    SCALE_GROUP_SIZE: tl.constexpr = 32
    K_packed = K_real // 2
    GRID_MN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)
    total_pids = GRID_MN * NUM_KSPLIT
    total_pids_padded = ((total_pids + 7) // 8) * 8

    pid_unified = tl.program_id(axis=0)
    pid_unified = _remap_xcd(pid_unified, total_pids_padded, NUM_XCDS=8)

    if pid_unified < total_pids:
        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) < K_real:
            num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE, BLOCK_SIZE_K)

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

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

            offs_bsn = (pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)) % N
            offs_ks_scale = (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_bsn[:, None] * stride_bsn + offs_ks_scale[None, :] * stride_bsk

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

            for k in tl.range(0, num_k_iter):
                a_bf16 = tl.load(a_ptrs).to(tl.float32)
                a_fp4, a_scales = _mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, QUANT_BLOCK)

                b_fp4 = tl.load(b_ptrs)

                b_scales = (
                    tl.load(b_scale_ptrs)
                    .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)
                )

                accumulator = tl.dot_scaled(a_fp4, a_scales, "e2m1", b_fp4, b_scales, "e2m1", accumulator)

                a_ptrs += BLOCK_SIZE_K * stride_ak
                b_ptrs += (BLOCK_SIZE_K // 2) * stride_bk
                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)


@triton.jit
def _fused_quant_shuffle_kernel(
    x_ptr, x_fp4_ptr, bs_ptr,
    stride_x_m, stride_x_n,
    stride_x_fp4_m, stride_x_fp4_n,
    M, N, scale_n_valid,
    SCALE_N: tl.constexpr,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    NUM_ITER: tl.constexpr,
    NUM_STAGES: tl.constexpr,
    MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
):
    pid_m = tl.program_id(0)
    start_n = tl.program_id(1) * NUM_ITER
    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE

    for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):
        x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
        x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
        x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
        x = tl.load(x_ptr + x_offs, mask=x_mask, other=0.0).to(tl.float32)

        out_tensor, bs_e8m0 = _mxfp4_quant_op(
            x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE
        )

        out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        out_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
        out_offs = out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
        out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
        tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)

        bs_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        bs_offs_n = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
        m_idx = bs_offs_m[:, None]
        n_idx = bs_offs_n[None, :]
        i0 = m_idx // 32
        i1 = (m_idx // 16) % 2
        i2 = m_idx % 16
        i3 = n_idx // 8
        i4 = (n_idx // 4) % 2
        i5 = n_idx % 4
        shuffled_offset = (i0 * (SCALE_N * 32) + i3 * 256 + i5 * 64 + i2 * 4 + i4 * 2 + i1)
        bs_valid = (bs_offs_m < M)[:, None] & (bs_offs_n < scale_n_valid)[None, :]
        bs_e8m0 = tl.where(bs_valid, bs_e8m0, 127)
        bs_store_mask = (m_idx < (M + 255) // 256 * 256) & (n_idx < SCALE_N)
        tl.store(bs_ptr + shuffled_offset, bs_e8m0, mask=bs_store_mask)


# ============================================================
# Cache: dict keyed by (M,K,N) + last-seen fast path
# Last-seen avoids dict lookup entirely for repeated calls
# ============================================================
_fused_cache = {}  # (M,K,N) -> config tuple
_asm_cache = {}    # (M,K,N) -> config tuple
_gemm_asm = None

# Last-seen fast path: avoids dict lookup for consecutive same-shape calls
_last_fused_key = None  # (M,K,N) tuple
_last_fused_cfg = None  # corresponding config
_last_asm_key = None
_last_asm_cfg = None


def _build_fused_config(M, K, N, A_device):
    """Build and cache fused config for a given (M,K,N). Called once per shape."""
    K_packed = K >> 1  # K // 2
    scale_n = (K + 31) >> 5  # (K + 31) // 32
    SCALE_N_B = ((scale_n + 7) >> 3) << 3  # round up to mult of 8

    BLOCK_SIZE_M = 16
    BLOCK_SIZE_N = 64 if M <= 16 else 128
    BLOCK_SIZE_K = 512

    num_pid_m = _triton_cdiv(M, BLOCK_SIZE_M)
    num_pid_n = _triton_cdiv(N, BLOCK_SIZE_N)
    base_blocks = num_pid_m * num_pid_n
    target_ksplit = max(1, 256 // max(1, base_blocks))

    NUM_KSPLIT = 1
    SPLITK_BLOCK_SIZE = K  # 2 * K_packed = K

    if target_ksplit > 1:
        sb, bk_adj, nk = _get_splitk_fn(K_packed, BLOCK_SIZE_K, target_ksplit)
        if bk_adj >= 512:
            BLOCK_SIZE_K = bk_adj
            SPLITK_BLOCK_SIZE = sb
            NUM_KSPLIT = nk

    if NUM_KSPLIT > 1:
        y_pp = _torch_empty((NUM_KSPLIT, M, N), dtype=_float32, device=A_device)
    else:
        y_pp = None
        SPLITK_BLOCK_SIZE = K  # 2 * K_packed

    y = _torch_empty((M, N), dtype=_bfloat16, device=A_device)

    total_blocks_raw = NUM_KSPLIT * num_pid_m * num_pid_n
    total_blocks = ((total_blocks_raw + 7) >> 3) << 3  # pad to mult of 8

    bs_stride_n = 32 * SCALE_N_B
    out_tensor = y if NUM_KSPLIT == 1 else y_pp

    # Pre-compute strides for output
    if NUM_KSPLIT == 1:
        stride_ck = 0
        stride_cm = y.stride(0)
        stride_cn = y.stride(1)
    else:
        stride_ck = y_pp.stride(0)
        stride_cm = y_pp.stride(1)
        stride_cn = y_pp.stride(2)

    # Pre-compute reduce params
    reduce_params = None
    if NUM_KSPLIT > 1:
        ACTUAL_KSPLIT = _triton_cdiv(K_packed, SPLITK_BLOCK_SIZE >> 1)
        grid_reduce = (_triton_cdiv(M, 16), _triton_cdiv(N, 64))
        np2_ksplit = _triton_np2(NUM_KSPLIT)
        reduce_params = (grid_reduce, ACTUAL_KSPLIT, np2_ksplit,
                         y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
                         y.stride(0), y.stride(1))

    # Pre-compute the grid tuple
    grid_tuple = (total_blocks,)

    return (M, N, K, BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K,
            NUM_KSPLIT, SPLITK_BLOCK_SIZE,
            y, y_pp, out_tensor, grid_tuple,
            stride_ck, stride_cm, stride_cn,
            bs_stride_n, reduce_params)


def _build_asm_config(M, K, N, A, B_shuffle, B_scale_sh):
    """Build and cache ASM config for a given (M,K,N). Called once per shape."""
    scale_n_valid = (K + 31) >> 5
    SCALE_M = ((M + 255) // 256) * 256
    SCALE_N = ((scale_n_valid + 7) >> 3) << 3
    padded_m = get_padded_m(M, N, K, 0)

    BSM = _triton_np2(M) if M <= 32 else 16
    BSN = 32
    NUM_ITER_Q = 2
    grid = (_triton_cdiv(M, BSM), _triton_cdiv(K, BSN * NUM_ITER_Q))

    ck_config = get_GEMM_config(M, N, K)
    kernel_name = ""
    split_k = 0
    if ck_config is not None:
        split_k = ck_config.get("splitK", 0) or 0
        kernel_name = ck_config["kernelName"]

    x_fp4 = _torch_empty((M, K >> 1), dtype=_uint8, device=A.device)
    bs_sh = _torch_full((SCALE_M, SCALE_N), 127, dtype=_uint8, device=A.device)
    out = _torch_empty((padded_m, N), dtype=_bfloat16, device=A.device)

    x_fp4_view = x_fp4.view(_fp4x2)
    bs_sh_view = bs_sh.view(_fp8_e8m0)
    out_view = out[:M] if M < padded_m else out

    stride_a0 = A.stride(0)
    stride_a1 = A.stride(1)
    stride_fp4_0 = x_fp4.stride(0)
    stride_fp4_1 = x_fp4.stride(1)

    return (scale_n_valid, SCALE_N, BSM, BSN, NUM_ITER_Q, grid,
            x_fp4, bs_sh, out, x_fp4_view, bs_sh_view, out_view,
            kernel_name, split_k,
            stride_a0, stride_a1, stride_fp4_0, stride_fp4_1)


# ============================================================
# The hot path: no warmup check, no dict lookup, minimal Python
# ============================================================
@torch.no_grad()
def _hot_fused(data):
    """Fused quant+GEMM hot path for M<=64. Absolute minimum Python overhead."""
    global _last_fused_key, _last_fused_cfg
    A, B, B_q, B_shuffle, B_scale_sh = data
    M = A.shape[0]
    K = A.shape[1]
    N = B_shuffle.shape[0]

    # Fast path: check last-seen key (avoids dict lookup for repeated shapes)
    key = (M, K, N)
    if key is _last_fused_key or key == _last_fused_key:
        c = _last_fused_cfg
    else:
        c = _fused_cache.get(key)
        if c is None:
            c = _build_fused_config(M, K, N, A.device)
            _fused_cache[key] = c
        _last_fused_key = key
        _last_fused_cfg = c

    # Unpack only what we need
    (_, _, _, BSM, BSN, BSK,
     NUM_KSPLIT, SPLITK_BLOCK_SIZE,
     y, y_pp, out_tensor, grid_tuple,
     stride_ck, stride_cm, stride_cn,
     bs_stride_n, reduce_params) = c

    # B views: must recompute since B changes between phases
    # .view() and .T are very cheap on contiguous tensors
    B_q_T = B_q.view(_uint8).T
    B_scale_u8 = B_scale_sh.view(_uint8)

    # Launch fused kernel - use pre-computed grid
    _fused_quant_gemm_kernel[grid_tuple](
        A, B_q_T, out_tensor, B_scale_u8,
        M, N, K,
        A.stride(0), A.stride(1),
        B_q_T.stride(0), B_q_T.stride(1),
        stride_ck, stride_cm, stride_cn,
        bs_stride_n, 1,
        BLOCK_SIZE_M=BSM,
        BLOCK_SIZE_N=BSN,
        BLOCK_SIZE_K=BSK,
        GROUP_SIZE_M=8,
        NUM_KSPLIT=NUM_KSPLIT,
        SPLITK_BLOCK_SIZE=SPLITK_BLOCK_SIZE,
        QUANT_BLOCK=32,
        num_warps=8,
        num_stages=2,
        waves_per_eu=0,
    )

    if reduce_params is not None:
        grid_reduce, ACTUAL_KSPLIT, np2_ksplit, s0, s1, s2, sy0, sy1 = reduce_params
        _reduce_kernel[grid_reduce](
            y_pp, y, M, N,
            s0, s1, s2,
            sy0, sy1,
            16, 64, ACTUAL_KSPLIT, np2_ksplit,
        )

    return y


@torch.no_grad()
def _hot_asm(data):
    """ASM GEMM hot path for M>64. Absolute minimum Python overhead."""
    global _last_asm_key, _last_asm_cfg
    A, B, B_q, B_shuffle, B_scale_sh = data
    M = A.shape[0]
    K = A.shape[1]
    N = B_shuffle.shape[0]

    key = (M, K, N)
    if key is _last_asm_key or key == _last_asm_key:
        c = _last_asm_cfg
    else:
        c = _asm_cache.get(key)
        if c is None:
            c = _build_asm_config(M, K, N, A, B_shuffle, B_scale_sh)
            _asm_cache[key] = c
        _last_asm_key = key
        _last_asm_cfg = c

    (scale_n_valid, SCALE_N, BSM, BSN, NUM_ITER_Q, grid,
     x_fp4, bs_sh, out, x_fp4_view, bs_sh_view, out_view,
     kernel_name, split_k,
     stride_a0, stride_a1, stride_fp4_0, stride_fp4_1) = c

    _fused_quant_shuffle_kernel[grid](
        A, x_fp4, bs_sh,
        stride_a0, stride_a1,
        stride_fp4_0, stride_fp4_1,
        M, K, scale_n_valid,
        SCALE_N=SCALE_N, BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=BSN,
        NUM_ITER=NUM_ITER_Q, NUM_STAGES=NUM_ITER_Q, MXFP4_QUANT_BLOCK_SIZE=32,
        num_warps=1, waves_per_eu=0, num_stages=NUM_ITER_Q,
    )

    _gemm_asm(x_fp4_view, B_shuffle, bs_sh_view, B_scale_sh,
              out, kernel_name, None, 1.0, 0.0, True, split_k)
    return out_view


@torch.no_grad()
def _hot_dispatch(data):
    """Dispatch to fused or ASM based on M. No warmup check."""
    M = data[0].shape[0]
    if M <= 64:
        return _hot_fused(data)
    else:
        return _hot_asm(data)


def _warmup_kernel(data):
    """Warmup path: initializes aiter, then swaps custom_kernel to hot path."""
    global custom_kernel, _gemm_asm

    A, B, B_q, B_shuffle, B_scale_sh = data
    M, K = A.shape
    N = B_shuffle.shape[0]

    scale_n_valid = (K + 31) >> 5
    SCALE_M = ((M + 255) // 256) * 256
    SCALE_N = ((scale_n_valid + 7) >> 3) << 3
    BSM = _triton_np2(M) if M <= 32 else 16
    grid = (_triton_cdiv(M, BSM), _triton_cdiv(K, 32))

    x_fp4 = _torch_empty((M, K >> 1), dtype=_uint8, device=A.device)
    bs_sh = _torch_full((SCALE_M, SCALE_N), 127, dtype=_uint8, device=A.device)

    _fused_quant_shuffle_kernel[grid](
        A, x_fp4, bs_sh,
        A.stride(0), A.stride(1),
        x_fp4.stride(0), x_fp4.stride(1),
        M, K, scale_n_valid,
        SCALE_N=SCALE_N, BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=32,
        NUM_ITER=1, NUM_STAGES=1, MXFP4_QUANT_BLOCK_SIZE=32,
        num_warps=1, waves_per_eu=0, num_stages=1,
    )

    result = aiter.gemm_a4w4(
        x_fp4.view(_fp4x2), B_shuffle,
        bs_sh.view(_fp8_e8m0), B_scale_sh,
        dtype=_bf16, bpreshuffle=True,
    )

    # Get ASM function
    try:
        _gemm_asm = torch.ops.aiter.gemm_a4w4_asm
    except Exception:
        try:
            import aiter.jit.core as _jc
            _gemm_asm = getattr(_jc, 'gemm_a4w4_asm', None)
        except Exception:
            pass

    # CRITICAL: swap custom_kernel to the hot path
    # This eliminates the warmup check from ALL future calls
    custom_kernel = _hot_dispatch

    return result


# Start with warmup - gets swapped to _hot_dispatch after first call
custom_kernel = _warmup_kernel
scrolls · 493 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 542580.

⋯ 1 unchanged lines
#!POPCORN gpu MI355X
"""
- v105: v99 + num_stages=3 for fused kernel (was 2). More pipeline stages
- for better overlap of memory loads with MFMA compute in the K loop.
+ v142: Extreme Python inline optimization.
+ - Function pointer swap: warmup function swaps to hot path (no if-check per call)
+ - List-indexed cache by M (no dict lookup in hot path)
+ - Pre-bound kernel function references (no global lookup)
+ - @torch.no_grad() to skip autograd overhead
+ - Minimal tuple unpacking in hot path
+ - Pre-computed B_q_T strides cached per shape
+ - All constants pre-computed during warmup/first-see
+ - ASM path for M>64 with same optimizations
"""
from task import input_t, output_t
⋯ 10 unchanged lines
_reduce_kernel = _gemm_mod._gemm_afp4wfp4_reduce_kernel
_get_splitk_fn = _gemm_mod.get_splitk
+ # Pre-bind dtype constants
_fp4x2 = dtypes.fp4x2
_fp8_e8m0 = dtypes.fp8_e8m0
_bf16 = dtypes.bf16
+ _uint8 = torch.uint8
+ _bfloat16 = torch.bfloat16
+ _float32 = torch.float32
+ # Pre-bind torch functions to avoid global lookups
+ _torch_empty = torch.empty
+ _torch_full = torch.full
+ _triton_cdiv = triton.cdiv
+ _triton_np2 = triton.next_power_of_2
+
@triton.jit
def _remap_xcd(pid, num_pids, NUM_XCDS: tl.constexpr):
chunk_size = tl.cdiv(num_pids, NUM_XCDS)
⋯ 65 unchanged lines
b_ptrs = b_ptr + offs_k_packed[:, None] * stride_bk + offs_bn[None, :] * stride_bn
offs_bsn = (pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)) % N
- offs_ks_scale = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE * 32)) + tl.arange(
+ offs_ks_scale = (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_bsn[:, None] * stride_bsn + offs_ks_scale[None, :] * stride_bsk
⋯ 78 unchanged lines
tl.store(bs_ptr + shuffled_offset, bs_e8m0, mask=bs_store_mask)
- _cache_asm = {}
- _cache_fused = {}
+ # ============================================================
+ # Cache: dict keyed by (M,K,N) + last-seen fast path
+ # Last-seen avoids dict lookup entirely for repeated calls
+ # ============================================================
+ _fused_cache = {} # (M,K,N) -> config tuple
+ _asm_cache = {} # (M,K,N) -> config tuple
_gemm_asm = None
- _warmup_done = False
+ # Last-seen fast path: avoids dict lookup for consecutive same-shape calls
+ _last_fused_key = None # (M,K,N) tuple
+ _last_fused_cfg = None # corresponding config
+ _last_asm_key = None
+ _last_asm_cfg = None
- def custom_kernel(data: input_t) -> output_t:
- global _gemm_asm, _warmup_done
- A, B, B_q, B_shuffle, B_scale_sh = data
- M, K = A.shape
- N = B_shuffle.shape[0]
+ def _build_fused_config(M, K, N, A_device):
+ """Build and cache fused config for a given (M,K,N). Called once per shape."""
+ K_packed = K >> 1 # K // 2
+ scale_n = (K + 31) >> 5 # (K + 31) // 32
+ SCALE_N_B = ((scale_n + 7) >> 3) << 3 # round up to mult of 8
- use_fused = (M <= 32)
+ BLOCK_SIZE_M = 16
+ BLOCK_SIZE_N = 64 if M <= 16 else 128
+ BLOCK_SIZE_K = 512
- if not _warmup_done:
- scale_n_valid = (K + 31) // 32
- SCALE_M = ((M + 255) // 256) * 256
- SCALE_N = ((scale_n_valid + 7) // 8) * 8
- BSM = triton.next_power_of_2(M) if M <= 32 else 16
- grid = (triton.cdiv(M, BSM), triton.cdiv(K, 32))
+ num_pid_m = _triton_cdiv(M, BLOCK_SIZE_M)
+ num_pid_n = _triton_cdiv(N, BLOCK_SIZE_N)
+ base_blocks = num_pid_m * num_pid_n
+ target_ksplit = max(1, 256 // max(1, base_blocks))
- x_fp4 = torch.empty((M, K // 2), dtype=torch.uint8, device=A.device)
- bs_sh = torch.full((SCALE_M, SCALE_N), 127, dtype=torch.uint8, device=A.device)
+ NUM_KSPLIT = 1
+ SPLITK_BLOCK_SIZE = K # 2 * K_packed = K
- _fused_quant_shuffle_kernel[grid](
- A, x_fp4, bs_sh,
- A.stride(0), A.stride(1),
- x_fp4.stride(0), x_fp4.stride(1),
- M, K, scale_n_valid,
- SCALE_N=SCALE_N, BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=32,
- NUM_ITER=1, NUM_STAGES=1, MXFP4_QUANT_BLOCK_SIZE=32,
- num_warps=1, waves_per_eu=0, num_stages=1,
- )
+ if target_ksplit > 1:
+ sb, bk_adj, nk = _get_splitk_fn(K_packed, BLOCK_SIZE_K, target_ksplit)
+ if bk_adj >= 512:
+ BLOCK_SIZE_K = bk_adj
+ SPLITK_BLOCK_SIZE = sb
+ NUM_KSPLIT = nk
- result = aiter.gemm_a4w4(
- x_fp4.view(_fp4x2), B_shuffle,
- bs_sh.view(_fp8_e8m0), B_scale_sh,
- dtype=_bf16, bpreshuffle=True,
- )
- _warmup_done = True
- try:
- _gemm_asm = torch.ops.aiter.gemm_a4w4_asm
- except Exception:
- try:
- import aiter.jit.core as _jc
- _gemm_asm = getattr(_jc, 'gemm_a4w4_asm', None)
- except Exception:
- pass
- return result
+ if NUM_KSPLIT > 1:
+ y_pp = _torch_empty((NUM_KSPLIT, M, N), dtype=_float32, device=A_device)
+ else:
+ y_pp = None
+ SPLITK_BLOCK_SIZE = K # 2 * K_packed
- if use_fused:
- key = (M, K, N)
- c = _cache_fused.get(key)
- if c is None:
- K_packed = K // 2
- scale_n = (K + 31) // 32
- SCALE_N_B = ((scale_n + 7) // 8) * 8
+ y = _torch_empty((M, N), dtype=_bfloat16, device=A_device)
- BLOCK_SIZE_M = 16
- BLOCK_SIZE_N = 64 if M <= 16 else 128
- BLOCK_SIZE_K = 512
+ total_blocks_raw = NUM_KSPLIT * num_pid_m * num_pid_n
+ total_blocks = ((total_blocks_raw + 7) >> 3) << 3 # pad to mult of 8
- base_blocks = triton.cdiv(M, BLOCK_SIZE_M) * triton.cdiv(N, BLOCK_SIZE_N)
- target_ksplit = max(1, 256 // max(1, base_blocks))
+ bs_stride_n = 32 * SCALE_N_B
+ out_tensor = y if NUM_KSPLIT == 1 else y_pp
- if target_ksplit > 1:
- SPLITK_BLOCK_SIZE, BLOCK_SIZE_K_adj, NUM_KSPLIT = _get_splitk_fn(
- K_packed, BLOCK_SIZE_K, target_ksplit
- )
- if BLOCK_SIZE_K_adj < 512:
- BLOCK_SIZE_K_adj = 512
- SPLITK_BLOCK_SIZE = 2 * K_packed
- NUM_KSPLIT = 1
- else:
- BLOCK_SIZE_K = BLOCK_SIZE_K_adj
- else:
- NUM_KSPLIT = 1
- SPLITK_BLOCK_SIZE = 2 * K_packed
+ # Pre-compute strides for output
+ if NUM_KSPLIT == 1:
+ stride_ck = 0
+ stride_cm = y.stride(0)
+ stride_cn = y.stride(1)
+ else:
+ stride_ck = y_pp.stride(0)
+ stride_cm = y_pp.stride(1)
+ stride_cn = y_pp.stride(2)
- if NUM_KSPLIT > 1:
- y_pp = torch.empty((NUM_KSPLIT, M, N), dtype=torch.float32, device=A.device)
- else:
- y_pp = None
- SPLITK_BLOCK_SIZE = 2 * K_packed
+ # Pre-compute reduce params
+ reduce_params = None
+ if NUM_KSPLIT > 1:
+ ACTUAL_KSPLIT = _triton_cdiv(K_packed, SPLITK_BLOCK_SIZE >> 1)
+ grid_reduce = (_triton_cdiv(M, 16), _triton_cdiv(N, 64))
+ np2_ksplit = _triton_np2(NUM_KSPLIT)
+ reduce_params = (grid_reduce, ACTUAL_KSPLIT, np2_ksplit,
+ y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
+ y.stride(0), y.stride(1))
- y = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
+ # Pre-compute the grid tuple
+ grid_tuple = (total_blocks,)
- total_blocks_raw = NUM_KSPLIT * triton.cdiv(M, BLOCK_SIZE_M) * triton.cdiv(N, BLOCK_SIZE_N)
- total_blocks = ((total_blocks_raw + 7) // 8) * 8
+ return (M, N, K, BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K,
+ NUM_KSPLIT, SPLITK_BLOCK_SIZE,
+ y, y_pp, out_tensor, grid_tuple,
+ stride_ck, stride_cm, stride_cn,
+ bs_stride_n, reduce_params)
- bs_stride_n = 32 * SCALE_N_B
- bs_stride_k = 1
- c = (K_packed, SCALE_N_B, BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K,
- NUM_KSPLIT, SPLITK_BLOCK_SIZE,
- y, y_pp, total_blocks, bs_stride_n, bs_stride_k)
- _cache_fused[key] = c
+ def _build_asm_config(M, K, N, A, B_shuffle, B_scale_sh):
+ """Build and cache ASM config for a given (M,K,N). Called once per shape."""
+ scale_n_valid = (K + 31) >> 5
+ SCALE_M = ((M + 255) // 256) * 256
+ SCALE_N = ((scale_n_valid + 7) >> 3) << 3
+ padded_m = get_padded_m(M, N, K, 0)
- (K_packed, SCALE_N_B, BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K,
- NUM_KSPLIT, SPLITK_BLOCK_SIZE,
- y, y_pp, total_blocks, bs_stride_n, bs_stride_k) = c
+ BSM = _triton_np2(M) if M <= 32 else 16
+ BSN = 32
+ NUM_ITER_Q = 2
+ grid = (_triton_cdiv(M, BSM), _triton_cdiv(K, BSN * NUM_ITER_Q))
- B_q_u8 = B_q.view(torch.uint8) if B_q.dtype != torch.uint8 else B_q
- B_q_T = B_q_u8.T
- B_scale_u8 = B_scale_sh.view(torch.uint8)
+ ck_config = get_GEMM_config(M, N, K)
+ kernel_name = ""
+ split_k = 0
+ if ck_config is not None:
+ split_k = ck_config.get("splitK", 0) or 0
+ kernel_name = ck_config["kernelName"]
- out_tensor = y if NUM_KSPLIT == 1 else y_pp
+ x_fp4 = _torch_empty((M, K >> 1), dtype=_uint8, device=A.device)
+ bs_sh = _torch_full((SCALE_M, SCALE_N), 127, dtype=_uint8, device=A.device)
+ out = _torch_empty((padded_m, N), dtype=_bfloat16, device=A.device)
- _fused_quant_gemm_kernel[(total_blocks,)](
- A, B_q_T, out_tensor, B_scale_u8,
- M, N, K,
- A.stride(0), A.stride(1),
- B_q_T.stride(0), B_q_T.stride(1),
- 0 if NUM_KSPLIT == 1 else y_pp.stride(0),
- y.stride(0) if NUM_KSPLIT == 1 else y_pp.stride(1),
- y.stride(1) if NUM_KSPLIT == 1 else y_pp.stride(2),
- bs_stride_n, bs_stride_k,
- BLOCK_SIZE_M=BLOCK_SIZE_M,
- BLOCK_SIZE_N=BLOCK_SIZE_N,
- BLOCK_SIZE_K=BLOCK_SIZE_K,
- GROUP_SIZE_M=8,
- NUM_KSPLIT=NUM_KSPLIT,
- SPLITK_BLOCK_SIZE=SPLITK_BLOCK_SIZE,
- QUANT_BLOCK=32,
- num_warps=8,
- num_stages=3, # KEY CHANGE: was 2
- waves_per_eu=0,
+ x_fp4_view = x_fp4.view(_fp4x2)
+ bs_sh_view = bs_sh.view(_fp8_e8m0)
+ out_view = out[:M] if M < padded_m else out
+
+ stride_a0 = A.stride(0)
+ stride_a1 = A.stride(1)
+ stride_fp4_0 = x_fp4.stride(0)
+ stride_fp4_1 = x_fp4.stride(1)
+
+ return (scale_n_valid, SCALE_N, BSM, BSN, NUM_ITER_Q, grid,
+ x_fp4, bs_sh, out, x_fp4_view, bs_sh_view, out_view,
+ kernel_name, split_k,
+ stride_a0, stride_a1, stride_fp4_0, stride_fp4_1)
+
+
+ # ============================================================
+ # The hot path: no warmup check, no dict lookup, minimal Python
+ # ============================================================
+ @torch.no_grad()
+ def _hot_fused(data):
+ """Fused quant+GEMM hot path for M<=64. Absolute minimum Python overhead."""
+ global _last_fused_key, _last_fused_cfg
+ A, B, B_q, B_shuffle, B_scale_sh = data
+ M = A.shape[0]
+ K = A.shape[1]
+ N = B_shuffle.shape[0]
+
+ # Fast path: check last-seen key (avoids dict lookup for repeated shapes)
+ key = (M, K, N)
+ if key is _last_fused_key or key == _last_fused_key:
+ c = _last_fused_cfg
+ else:
+ c = _fused_cache.get(key)
+ if c is None:
+ c = _build_fused_config(M, K, N, A.device)
+ _fused_cache[key] = c
+ _last_fused_key = key
+ _last_fused_cfg = c
+
+ # Unpack only what we need
+ (_, _, _, BSM, BSN, BSK,
+ NUM_KSPLIT, SPLITK_BLOCK_SIZE,
+ y, y_pp, out_tensor, grid_tuple,
+ stride_ck, stride_cm, stride_cn,
+ bs_stride_n, reduce_params) = c
+
+ # B views: must recompute since B changes between phases
+ # .view() and .T are very cheap on contiguous tensors
+ B_q_T = B_q.view(_uint8).T
+ B_scale_u8 = B_scale_sh.view(_uint8)
+
+ # Launch fused kernel - use pre-computed grid
+ _fused_quant_gemm_kernel[grid_tuple](
+ A, B_q_T, out_tensor, B_scale_u8,
+ M, N, K,
+ A.stride(0), A.stride(1),
+ B_q_T.stride(0), B_q_T.stride(1),
+ stride_ck, stride_cm, stride_cn,
+ bs_stride_n, 1,
+ BLOCK_SIZE_M=BSM,
+ BLOCK_SIZE_N=BSN,
+ BLOCK_SIZE_K=BSK,
+ GROUP_SIZE_M=8,
+ NUM_KSPLIT=NUM_KSPLIT,
+ SPLITK_BLOCK_SIZE=SPLITK_BLOCK_SIZE,
+ QUANT_BLOCK=32,
+ num_warps=8,
+ num_stages=2,
+ waves_per_eu=0,
+ )
+
+ if reduce_params is not None:
+ grid_reduce, ACTUAL_KSPLIT, np2_ksplit, s0, s1, s2, sy0, sy1 = reduce_params
+ _reduce_kernel[grid_reduce](
+ y_pp, y, M, N,
+ s0, s1, s2,
+ sy0, sy1,
+ 16, 64, ACTUAL_KSPLIT, np2_ksplit,
)
- if NUM_KSPLIT > 1:
- ACTUAL_KSPLIT = triton.cdiv(K_packed, (SPLITK_BLOCK_SIZE // 2))
- grid_reduce = (triton.cdiv(M, 16), triton.cdiv(N, 64))
- _reduce_kernel[grid_reduce](
- 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
- return y
+ @torch.no_grad()
+ def _hot_asm(data):
+ """ASM GEMM hot path for M>64. Absolute minimum Python overhead."""
+ global _last_asm_key, _last_asm_cfg
+ A, B, B_q, B_shuffle, B_scale_sh = data
+ M = A.shape[0]
+ K = A.shape[1]
+ N = B_shuffle.shape[0]
+
+ key = (M, K, N)
+ if key is _last_asm_key or key == _last_asm_key:
+ c = _last_asm_cfg
else:
- key = (M, K, N)
- c = _cache_asm.get(key)
+ c = _asm_cache.get(key)
if c is None:
- scale_n_valid = (K + 31) // 32
- SCALE_M = ((M + 255) // 256) * 256
- SCALE_N = ((scale_n_valid + 7) // 8) * 8
- padded_m = get_padded_m(M, N, K, 0)
+ c = _build_asm_config(M, K, N, A, B_shuffle, B_scale_sh)
+ _asm_cache[key] = c
+ _last_asm_key = key
+ _last_asm_cfg = c
- BSM = 16
- BSN = 32
- NUM_ITER_Q = 2
- grid = (triton.cdiv(M, BSM), triton.cdiv(K, BSN * NUM_ITER_Q))
+ (scale_n_valid, SCALE_N, BSM, BSN, NUM_ITER_Q, grid,
+ x_fp4, bs_sh, out, x_fp4_view, bs_sh_view, out_view,
+ kernel_name, split_k,
+ stride_a0, stride_a1, stride_fp4_0, stride_fp4_1) = c
- ck_config = get_GEMM_config(M, N, K)
- kernel_name = ""
- split_k = 0
- if ck_config is not None:
- split_k = ck_config.get("splitK", 0) or 0
- kernel_name = ck_config["kernelName"]
+ _fused_quant_shuffle_kernel[grid](
+ A, x_fp4, bs_sh,
+ stride_a0, stride_a1,
+ stride_fp4_0, stride_fp4_1,
+ M, K, scale_n_valid,
+ SCALE_N=SCALE_N, BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=BSN,
+ NUM_ITER=NUM_ITER_Q, NUM_STAGES=NUM_ITER_Q, MXFP4_QUANT_BLOCK_SIZE=32,
+ num_warps=1, waves_per_eu=0, num_stages=NUM_ITER_Q,
+ )
- x_fp4 = torch.empty((M, K // 2), dtype=torch.uint8, device=A.device)
- bs_sh = torch.full((SCALE_M, SCALE_N), 127, dtype=torch.uint8, device=A.device)
- out = torch.empty((padded_m, N), dtype=torch.bfloat16, device=A.device)
+ _gemm_asm(x_fp4_view, B_shuffle, bs_sh_view, B_scale_sh,
+ out, kernel_name, None, 1.0, 0.0, True, split_k)
+ return out_view
- x_fp4_view = x_fp4.view(_fp4x2)
- bs_sh_view = bs_sh.view(_fp8_e8m0)
- out_view = out[:M] if M < padded_m else out
- c = (scale_n_valid, SCALE_N, BSM, BSN, NUM_ITER_Q, grid,
- x_fp4, bs_sh, out, x_fp4_view, bs_sh_view, out_view,
- kernel_name, split_k,
- A.stride(0), A.stride(1), x_fp4.stride(0), x_fp4.stride(1))
- _cache_asm[key] = c
+ @torch.no_grad()
+ def _hot_dispatch(data):
+ """Dispatch to fused or ASM based on M. No warmup check."""
+ M = data[0].shape[0]
+ if M <= 64:
+ return _hot_fused(data)
+ else:
+ return _hot_asm(data)
- (scale_n_valid, SCALE_N, BSM, BSN, NUM_ITER_Q, grid,
- x_fp4, bs_sh, out, x_fp4_view, bs_sh_view, out_view,
- kernel_name, split_k,
- stride_a0, stride_a1, stride_fp4_0, stride_fp4_1) = c
- _fused_quant_shuffle_kernel[grid](
- A, x_fp4, bs_sh,
- stride_a0, stride_a1,
- stride_fp4_0, stride_fp4_1,
- M, K, scale_n_valid,
- SCALE_N=SCALE_N, BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=BSN,
- NUM_ITER=NUM_ITER_Q, NUM_STAGES=NUM_ITER_Q, MXFP4_QUANT_BLOCK_SIZE=32,
- num_warps=1, waves_per_eu=0, num_stages=NUM_ITER_Q,
- )
+ def _warmup_kernel(data):
+ """Warmup path: initializes aiter, then swaps custom_kernel to hot path."""
+ global custom_kernel, _gemm_asm
- if _gemm_asm is not None:
- _gemm_asm(x_fp4_view, B_shuffle, bs_sh_view, B_scale_sh,
- out, kernel_name, None, 1.0, 0.0, True, split_k)
- return out_view
+ A, B, B_q, B_shuffle, B_scale_sh = data
+ M, K = A.shape
+ N = B_shuffle.shape[0]
- return aiter.gemm_a4w4(
- x_fp4_view, B_shuffle, bs_sh_view, B_scale_sh,
- dtype=_bf16, bpreshuffle=True,
- )
+ scale_n_valid = (K + 31) >> 5
+ SCALE_M = ((M + 255) // 256) * 256
+ SCALE_N = ((scale_n_valid + 7) >> 3) << 3
+ BSM = _triton_np2(M) if M <= 32 else 16
+ grid = (_triton_cdiv(M, BSM), _triton_cdiv(K, 32))
+
+ x_fp4 = _torch_empty((M, K >> 1), dtype=_uint8, device=A.device)
+ bs_sh = _torch_full((SCALE_M, SCALE_N), 127, dtype=_uint8, device=A.device)
+
+ _fused_quant_shuffle_kernel[grid](
+ A, x_fp4, bs_sh,
+ A.stride(0), A.stride(1),
+ x_fp4.stride(0), x_fp4.stride(1),
+ M, K, scale_n_valid,
+ SCALE_N=SCALE_N, BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=32,
+ NUM_ITER=1, NUM_STAGES=1, MXFP4_QUANT_BLOCK_SIZE=32,
+ num_warps=1, waves_per_eu=0, num_stages=1,
+ )
+
+ result = aiter.gemm_a4w4(
+ x_fp4.view(_fp4x2), B_shuffle,
+ bs_sh.view(_fp8_e8m0), B_scale_sh,
+ dtype=_bf16, bpreshuffle=True,
+ )
+
+ # Get ASM function
+ try:
+ _gemm_asm = torch.ops.aiter.gemm_a4w4_asm
+ except Exception:
+ try:
+ import aiter.jit.core as _jc
+ _gemm_asm = getattr(_jc, 'gemm_a4w4_asm', None)
+ except Exception:
+ pass
+
+ # CRITICAL: swap custom_kernel to the hot path
+ # This eliminates the warmup check from ALL future calls
+ custom_kernel = _hot_dispatch
+
+ return result
+
+
+ # Start with warmup - gets swapped to _hot_dispatch after first call
+ custom_kernel = _warmup_kernel
scrolls · 510 diff lines total

Best evidence level for this revision: reported

JSON