Skip to content
KernelIndex
Search⌘K

submission 561029

kkosey · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-561029?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
9.01µs
#107 of 1143
2026-03-15

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:abb28e16426ebc4332d35e0d1a9ae9dd8b67d787a573e8e138e431ed72ab5bb8
license declaredunknown
license concludedunknown
authorskkosey
imported2026-08-15

Techniques

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

fp4Optimized MXFP4 GEMM: Custom Triton kernel for small-K + ASM for large-K.
num-warps = 1NUM_WARPS = 1
split-kfrom aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
stages = 1num_stages = 1
tile-m = 32BLOCK_SIZE_M = 32
tile-n = 32BLOCK_SIZE_N = 32

Kernel source

submission.py544 lines
"""
Optimized MXFP4 GEMM: Custom Triton kernel for small-K + ASM for large-K.
The Triton kernel fuses bf16→MXFP4 quant + GEMM + shuffled B_scale read.
MI355X: 256 CUs, 8 XCDs, gfx950, native FP4 MFMA 16x16.
"""
from task import input_t, output_t
import torch
import triton
import triton.language as tl
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid, remap_xcd
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import _get_config
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk

try:
    from aiter.ops.gemm_op_a4w4 import get_padded_m as _get_padded_m
except ImportError:
    _get_padded_m = None

# Raw ASM dispatch: bypass torch.ops.aiter wrapper overhead (~0.4-0.6µs savings)
_raw_gemm_fn = None
def _init_raw_gemm():
    global _raw_gemm_fn
    if _raw_gemm_fn is not None:
        return
    try:
        from aiter.jit.core import get_module
        mod = get_module("module_gemm_a4w4_asm")
        _raw_gemm_fn = mod.gemm_a4w4_asm
    except Exception:
        pass  # .so not built yet, will use wrapper

_fp4x2 = dtypes.fp4x2
_fp8_e8m0 = dtypes.fp8_e8m0
_bf16 = dtypes.bf16

def _asm_kernel_name(tile_m, tile_n):
    name = f"f4gemm_bf16_per1x32Fp4_BpreShuffle_{tile_m}x{tile_n}"
    return f"_ZN5aiter{len(name)}{name}E"

_DEFAULT_32x128 = _asm_kernel_name(32, 128)


# ---------------------------------------------------------------------------
# Custom GEMM kernel: bf16 A × fp4 B → bf16 C, reading SHUFFLED B_scale
# Based on aiter's _gemm_a16wfp4_kernel but with inline e8m0 unshuffle.
# Eliminates the need for a separate unshuffle kernel launch.
# ---------------------------------------------------------------------------
@triton.jit
def _gemm_a16wfp4_shuffled_scale_kernel(
    a_ptr, b_ptr, c_ptr, b_scales_ptr,
    M, N, K,  # K = K_half (packed)
    stride_am, stride_ak,
    stride_bk, stride_bn,
    stride_cm, stride_cn,
    stride_ck,  # 0 for no splitK, M*N for splitK (fp32 output)
    # constexpr meta-parameters
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    BLOCK_SIZE_K: tl.constexpr,
    GROUP_SIZE_M: tl.constexpr,
    SN_SCALE: tl.constexpr,  # padded scale columns (sn from e8m0_shuffle)
    EVEN_K: tl.constexpr,
    NUM_KSPLIT: tl.constexpr,
    SPLITK_BLOCK_SIZE: tl.constexpr,
    num_warps: tl.constexpr,
    num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr,
    matrix_instr_nonkdim: tl.constexpr,
    cache_modifier: tl.constexpr,
):
    """GEMM C = A @ B^T with inline MXFP4 quant of A and shuffled B_scale read."""
    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)

    SCALE_GROUP_SIZE: tl.constexpr = 32

    GRID_MN = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)
    pid_unified = tl.program_id(axis=0)

    if NUM_KSPLIT > 1:
        pid_unified = remap_xcd(pid_unified, GRID_MN * NUM_KSPLIT, NUM_XCDS=8)
        pid_k = pid_unified % NUM_KSPLIT
        pid = pid_unified // NUM_KSPLIT
    else:
        pid_k = 0
        pid = pid_unified

    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)
    tl.assume(pid_k >= 0)

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

        # A pointers (bf16) — offset by splitK range
        offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
        k_offset_bf16 = pid_k * SPLITK_BLOCK_SIZE  # bf16 element offset
        offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
        a_ptrs = a_ptr + offs_am[:, None] * stride_am + (k_offset_bf16 + offs_k_bf16[None, :]) * stride_ak

        # B pointers (fp4 packed as uint8) — offset by splitK range
        offs_k = tl.arange(0, BLOCK_SIZE_K // 2)
        k_offset_packed = pid_k * (SPLITK_BLOCK_SIZE // 2)  # packed byte offset
        offs_bn = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
        b_ptrs = b_ptr + (k_offset_packed + offs_k[:, None]) * stride_bk + offs_bn[None, :] * stride_bn

        # Pre-compute row decomposition for shuffled B_scale access
        bs_d0 = (offs_bn // 32)[:, None]
        bs_d1 = ((offs_bn % 32) // 16)[:, None]
        bs_d2 = (offs_bn % 16)[:, None]

        # Scale column tracking — start at splitK offset
        SCALES_PER_BLOCK: tl.constexpr = BLOCK_SIZE_K // SCALE_GROUP_SIZE
        cur_scale_col = pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE)
        scale_k_range = tl.arange(0, SCALES_PER_BLOCK)

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

        for k in range(0, num_k_iter):
            # Load B_scale from SHUFFLED layout (inline unshuffle)
            cols = cur_scale_col + scale_k_range
            bs_d3 = (cols // 8)[None, :]
            bs_d4 = ((cols % 8) // 4)[None, :]
            bs_d5 = (cols % 4)[None, :]
            shuffled_idx = (bs_d0 * (32 * SN_SCALE) + bs_d3 * 256
                            + bs_d5 * 64 + bs_d2 * 4 + bs_d4 * 2 + bs_d1)
            b_scales = tl.load(b_scales_ptr + shuffled_idx)

            # Load A (bf16) and B (fp4 packed)
            if EVEN_K:
                a_bf16 = tl.load(a_ptrs)
                b = tl.load(b_ptrs, cache_modifier=cache_modifier)
            else:
                a_bf16 = tl.load(
                    a_ptrs,
                    mask=offs_k_bf16[None, :] < 2 * K - (pid_k * num_k_iter + k) * BLOCK_SIZE_K,
                    other=0,
                )
                b = tl.load(
                    b_ptrs,
                    mask=offs_k[:, None] < K - (pid_k * num_k_iter + k) * (BLOCK_SIZE_K // 2),
                    other=0,
                    cache_modifier=cache_modifier,
                )

            # In-register quant: bf16 A → mxfp4 + e8m0 scale
            a, a_scales = _mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)

            # Scaled dot product
            accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")

            # Advance pointers
            a_ptrs += BLOCK_SIZE_K * stride_ak
            b_ptrs += (BLOCK_SIZE_K // 2) * stride_bk
            cur_scale_col += SCALES_PER_BLOCK

        # Store output
        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)
        if NUM_KSPLIT == 1:
            c = accumulator.to(c_ptr.type.element_ty)
        else:
            c = accumulator  # keep fp32 for splitK
        tl.store(c_ptrs, c, mask=c_mask)


@triton.jit
def _splitk_reduce_kernel(
    y_pp_ptr, y_ptr,
    M, N, NUM_KSPLIT: tl.constexpr,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
    num_warps: tl.constexpr,
):
    """Reduce splitK partial results: y = sum(y_pp[k]) converted to bf16."""
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)
    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
    base = offs_m[:, None] * N + offs_n[None, :]
    for k in range(NUM_KSPLIT):
        val = tl.load(y_pp_ptr + k * M * N + base, mask=mask, other=0.0)
        acc += val
    tl.store(y_ptr + base, acc.to(tl.bfloat16), mask=mask)


def _use_triton(M, K):
    return (K <= 1024 and M <= 64) or M <= 16


_buf = {}


@triton.jit
def _fused_quant_shuffle_kernel(
    x_ptr,
    x_fp4_ptr,
    shuffled_bs_ptr,
    stride_x_m_in,
    stride_x_n_in,
    stride_x_fp4_m_in,
    stride_x_fp4_n_in,
    M: tl.constexpr,
    N: tl.constexpr,
    sn: 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,
    EVEN_M_N: tl.constexpr,
):
    pid_m = tl.program_id(0)
    start_n = tl.program_id(1) * NUM_ITER
    stride_x_m = tl.cast(stride_x_m_in, tl.int64)
    stride_x_n = tl.cast(stride_x_n_in, tl.int64)
    stride_x_fp4_m = tl.cast(stride_x_fp4_m_in, tl.int64)
    stride_x_fp4_n = tl.cast(stride_x_fp4_n_in, tl.int64)

    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

        if EVEN_M_N:
            x = tl.load(x_ptr + x_offs).to(tl.float32)
        else:
            x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
            x = tl.load(x_ptr + x_offs, mask=x_mask).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
        )

        if EVEN_M_N:
            tl.store(x_fp4_ptr + out_offs, out_tensor)
        else:
            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)

        row = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        col = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)

        d0 = (row // 32)[:, None]
        d1 = ((row % 32) // 16)[:, None]
        d2 = (row % 16)[:, None]
        d3 = (col // 8)[None, :]
        d4 = ((col % 8) // 4)[None, :]
        d5 = (col % 4)[None, :]

        flat_out = d0 * (32 * sn) + d3 * 256 + d5 * 64 + d2 * 4 + d4 * 2 + d1

        if EVEN_M_N:
            tl.store(shuffled_bs_ptr + flat_out, bs_e8m0)
        else:
            SCALE_N: tl.constexpr = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
            bs_mask = (row[:, None] < M) & (col[None, :] < SCALE_N)
            tl.store(shuffled_bs_ptr + flat_out, bs_e8m0, mask=bs_mask)


def _prepare_triton(M, N, K, device):
    K_half = K // 2
    SCALE_K = K // 32
    sn = (SCALE_K + 7) // 8 * 8  # padded scale cols (for shuffle formula)

    # Get config (same as wrapper: _get_config(M, N, K_half))
    raw_config, _ = _get_config(M, N, K_half)

    BSK = raw_config["BLOCK_SIZE_K"]
    if BSK >= 2 * K_half:
        BSK = triton.next_power_of_2(2 * K_half)
    BSK = max(BSK, 64)

    BSM = raw_config["BLOCK_SIZE_M"]
    BSN = raw_config["BLOCK_SIZE_N"]
    NW = raw_config["num_warps"]

    # Determine splitK based on shape
    grid_mn = triton.cdiv(M, BSM) * triton.cdiv(N, BSN)
    if K > 1024 and grid_mn < 128:
        # Large K, few MN blocks → use splitK
        # Use BSM that matches M for maximum MFMA utilization
        BSM = triton.next_power_of_2(min(M, 16))
        BSN = 64  # larger N-tiles for better data reuse
        BSK = 512  # larger K-tiles: fewer iterations, better register reuse
        NW = 4
        num_stages = 1
        grid_mn = triton.cdiv(M, BSM) * triton.cdiv(N, BSN)
        # Target ~2 waves (fewer tiles = larger work per tile = better efficiency)
        target_blocks = 256 * 2
        NUM_KSPLIT = max(1, min(target_blocks // max(grid_mn, 1), K_half // (BSK // 2)))
        # Use get_splitk to ensure EVEN_K
        SPLITK_BS, BSK, NUM_KSPLIT = get_splitk(K_half, BSK, NUM_KSPLIT)
        EVEN_K = (K_half % (BSK // 2) == 0) and (SPLITK_BS % BSK == 0) and (K_half % (SPLITK_BS // 2) == 0)
    else:
        NUM_KSPLIT = 1
        # Override: use BSN=64 for more grid parallelism
        if BSN > 64:
            BSN = 64
            NW = max(2, NW // 2)
        SPLITK_BS = 2 * K_half
        EVEN_K = (K_half % (BSK // 2) == 0) and (SPLITK_BS % BSK == 0) and (K_half % (SPLITK_BS // 2) == 0)
        num_stages = raw_config["num_stages"]

    # Filter config to only keys our custom kernel accepts
    config = {
        "BLOCK_SIZE_M": BSM,
        "BLOCK_SIZE_N": BSN,
        "BLOCK_SIZE_K": BSK,
        "GROUP_SIZE_M": raw_config["GROUP_SIZE_M"],
        "num_warps": NW,
        "num_stages": num_stages,
        "waves_per_eu": 0 if NUM_KSPLIT > 1 else raw_config["waves_per_eu"],
        "matrix_instr_nonkdim": raw_config["matrix_instr_nonkdim"],
        "cache_modifier": raw_config.get("cache_modifier", ".cg"),
        "NUM_KSPLIT": NUM_KSPLIT,
        "SPLITK_BLOCK_SIZE": SPLITK_BS,
    }

    gemm_out = torch.empty((M, N), dtype=_bf16, device=device)
    grid_size = NUM_KSPLIT * triton.cdiv(M, config["BLOCK_SIZE_M"]) * triton.cdiv(N, config["BLOCK_SIZE_N"])

    # Allocate splitK intermediate buffer if needed
    if NUM_KSPLIT > 1:
        y_pp = torch.empty((NUM_KSPLIT, M, N), dtype=torch.float32, device=device)
        reduce_bm = min(M, 16)
        reduce_bn = 64
        reduce_grid = (triton.cdiv(M, reduce_bm), triton.cdiv(N, reduce_bn))
    else:
        y_pp = None
        reduce_bm = 0
        reduce_bn = 0
        reduce_grid = None

    # Pre-compute all stride/scalar values as Python ints for fast hot path
    stride_am = K  # bf16 elements per row
    stride_ak = 1
    stride_bk = 1  # B_q is K-contiguous (after .T)
    stride_bn = K_half  # B_q row stride (after .T)
    stride_cm = N
    stride_cn = 1
    stride_ck = M * N if NUM_KSPLIT > 1 else 0

    return (gemm_out, config, sn, EVEN_K, grid_size, K_half, y_pp,
            NUM_KSPLIT, reduce_grid, reduce_bm, reduce_bn,
            stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_cn, stride_ck)


def _prepare_asm(M, N, K, device):
    # Cache raw ASM function on first ASM call
    _init_raw_gemm()

    MXFP4_QUANT_BLOCK_SIZE = 32
    SCALE_N = (K + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
    sm = (M + 255) // 256 * 256
    sn = (SCALE_N + 7) // 8 * 8

    tile_m = 32
    if _get_padded_m is not None:
        padded_m = _get_padded_m(M, N, K, tile_m)
    else:
        padded_m = ((M + tile_m - 1) // tile_m) * tile_m

    x_fp4 = torch.empty((padded_m, K // 2), dtype=torch.uint8, device=device)
    shuffled_scale = torch.zeros(sm * sn, dtype=torch.uint8, device=device)

    if M <= 32:
        NUM_ITER = 1
        BLOCK_SIZE_M = triton.next_power_of_2(M)
        BLOCK_SIZE_N = 32
        NUM_WARPS = 1
        NUM_STAGES = 1
    else:
        NUM_ITER = 1
        BLOCK_SIZE_M = 32
        BLOCK_SIZE_N = 64  # reduced from 128 for better CU utilization
        NUM_WARPS = 2
        NUM_STAGES = 1
    if K <= 1024:
        NUM_ITER = 1
        NUM_STAGES = 1
        BLOCK_SIZE_N = 32
        NUM_WARPS = 1
        BLOCK_SIZE_M = min(32, triton.next_power_of_2(M))

    EVEN_M_N = (M % BLOCK_SIZE_M == 0) and (K % (BLOCK_SIZE_N * NUM_ITER) == 0)
    grid = (
        triton.cdiv(M, BLOCK_SIZE_M),
        triton.cdiv(K, BLOCK_SIZE_N * NUM_ITER),
    )

    gemm_out = torch.empty((padded_m, N), dtype=_bf16, device=device)
    fp4_stride = x_fp4.stride()
    kernel_name = _DEFAULT_32x128

    # Determine splitK based on tile count (32×128 tiles)
    tiles_m = (padded_m + 31) // 32
    tiles_n = (N + 127) // 128
    grid_mn = tiles_m * tiles_n
    # Avoid log2_k_split=1 (known broken on MI355X)
    # Only use splitK for large K with few tiles
    if grid_mn < 32 and K >= 4096:
        log2_k_split = 4  # 16-way K-split for very few tiles + large K
    elif grid_mn < 128 and K >= 1024:
        log2_k_split = 3  # 8-way K-split for low tile counts
    else:
        log2_k_split = 0  # no split

    return (
        x_fp4, shuffled_scale, grid,
        BLOCK_SIZE_M, BLOCK_SIZE_N, NUM_ITER, NUM_WARPS, NUM_STAGES,
        EVEN_M_N, sn,
        gemm_out, kernel_name, log2_k_split,
        x_fp4.view(_fp4x2), shuffled_scale.view(sm, sn).view(_fp8_e8m0),
        gemm_out[:M], fp4_stride[0], fp4_stride[1], K, M,
    )


def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    M, K = A.shape
    N = B.shape[0]

    key = (M, K, N)
    use_t = _use_triton(M, K)

    if key not in _buf:
        if use_t:
            _buf[key] = ('t',) + _prepare_triton(M, N, K, A.device)
        else:
            _buf[key] = ('a',) + _prepare_asm(M, N, K, A.device)

    buf = _buf[key]

    if buf[0] == 't':
        (_, gemm_out, config, sn, EVEN_K, grid_size, K_half, y_pp,
         NUM_KSPLIT, reduce_grid, reduce_bm, reduce_bn,
         stride_am, stride_ak, stride_bk, stride_bn, stride_cm, stride_cn, stride_ck) = buf

        # Recompute B views each call (cheap views, avoids stale B caching)
        # Skip .T — strides are pre-cached, data_ptr is the same
        b_uint8 = B_q.view(torch.uint8)
        b_scale_raw = B_scale_sh.view(torch.uint8)

        if NUM_KSPLIT > 1:
            # SplitK: write fp32 partials to y_pp, then reduce
            _gemm_a16wfp4_shuffled_scale_kernel[(grid_size,)](
                A, b_uint8, y_pp, b_scale_raw,
                M, N, K_half,
                stride_am, stride_ak,
                stride_bk, stride_bn,
                stride_cm, stride_cn,
                stride_ck,
                SN_SCALE=sn,
                EVEN_K=EVEN_K,
                **config,
            )
            # Reduce: sum splitK partials → bf16 output
            _splitk_reduce_kernel[reduce_grid](
                y_pp, gemm_out, M, N,
                NUM_KSPLIT=NUM_KSPLIT,
                BLOCK_M=reduce_bm, BLOCK_N=reduce_bn,
                num_warps=4,
            )
        else:
            # No splitK: direct bf16 output
            _gemm_a16wfp4_shuffled_scale_kernel[(grid_size,)](
                A, b_uint8, gemm_out, b_scale_raw,
                M, N, K_half,
                stride_am, stride_ak,
                stride_bk, stride_bn,
                stride_cm, stride_cn,
                0,
                SN_SCALE=sn,
                EVEN_K=EVEN_K,
                **config,
            )
        return gemm_out
    else:
        (_, x_fp4, shuffled_scale, grid,
         BLOCK_SIZE_M, BLOCK_SIZE_N, NUM_ITER, NUM_WARPS, NUM_STAGES,
         EVEN_M_N, sn,
         gemm_out, kernel_name, log2_k_split,
         A_q, A_scale, out_view, fp4_s0, fp4_s1, K_val, M_val) = buf

        _fused_quant_shuffle_kernel[grid](
            A, x_fp4, shuffled_scale,
            K_val, 1, fp4_s0, fp4_s1,
            M=M_val, N=K_val, sn=sn,
            MXFP4_QUANT_BLOCK_SIZE=32,
            NUM_ITER=NUM_ITER,
            BLOCK_SIZE_M=BLOCK_SIZE_M,
            BLOCK_SIZE_N=BLOCK_SIZE_N,
            NUM_STAGES=NUM_STAGES,
            EVEN_M_N=EVEN_M_N,
            num_warps=NUM_WARPS,
            waves_per_eu=0, num_stages=1,
        )

        if _raw_gemm_fn is not None:
            _raw_gemm_fn(
                A_q, B_shuffle, A_scale, B_scale_sh,
                gemm_out, kernel_name,
                None, 1.0, 0.0, True, log2_k_split,
            )
        else:
            gemm_a4w4_asm(
                A_q, B_shuffle, A_scale, B_scale_sh,
                gemm_out, kernel_name,
                bpreshuffle=True, log2_k_split=log2_k_split,
            )
            _init_raw_gemm()  # Cache for subsequent calls

        return out_view
scrolls · 544 lines total

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

Best evidence level for this revision: reported

JSON