Skip to content
KernelIndex
Search⌘K

submission 713051

sean_nobricks · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:612e96a28636d4c07fb896e67ba319d8b314f298143b7e102aceef8a4c4b893f
license declaredunknown
license concludedunknown
authorssean_nobricks
imported2026-08-15

Techniques

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

autotunetriton.Config({'BLOCK_M': 16, 'BLOCK_N': 32, 'BLOCK_K': 256}, num_warps=4, num_stages=2),
fp4"""MXFP4 GEMM with shape-specialized dispatch.
num-warps = 4triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32, 'BLOCK_K': 256}, num_warps=4, num_stages=2),
split-kvery-small-`M` corner uses a split-K workspace reduction, while the rest of the
stages = 2triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32, 'BLOCK_K': 256}, num_warps=4, num_stages=2),
tile-k = 32BLOCK_K = 32
tile-m = 16BLOCK_M = 16
tile-n = 32BLOCK_M=16, BLOCK_N=32, BLOCK_K=256, SPLIT_K=SPLIT_K,

Kernel source

submission.py693 lines
"""MXFP4 GEMM with shape-specialized dispatch.

Path 1 (`K <= 512`) uses a fused Triton kernel that quantizes `A` in-register
and calls `tl.dot_scaled`.

Path 2 (`K > 512, M > 32`) quantizes `A` with Triton and calls the direct AITER
FP4 GEMM on preshuffled `B`.

Path 3 (`K > 512, M <= 32`) keeps quantization and GEMM fused in Triton. The
very-small-`M` corner uses a split-K workspace reduction, while the rest of the
small-`M` range uses direct split-K accumulation.
"""

import torch
import triton
import triton.language as tl
from task import input_t, output_t


# =============================================================================
# Software MXFP4 quant — used for K <= 512 fused path (Path 1)
# =============================================================================

@triton.jit
def _mxfp4_quant_tile(x, BLOCK_M: tl.constexpr, BLOCK_K: tl.constexpr):
    """Quantize (BLOCK_M, BLOCK_K) fp32 tile to MXFP4 in-register."""
    SG: tl.constexpr = 32
    NG: tl.constexpr = BLOCK_K // SG

    x = x.reshape(BLOCK_M, NG, SG)

    amax = tl.max(tl.abs(x), axis=2, keep_dims=True)
    amax_i = amax.to(tl.int32, bitcast=True)
    amax_i = (amax_i + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
    amax = amax_i.to(tl.float32, bitcast=True)
    scale_ub = tl.log2(amax).floor() - 2.0
    scale_ub = tl.clamp(scale_ub, min=-127.0, max=127.0)
    scales = scale_ub.to(tl.uint8) + 127

    qx = x * tl.exp2(-scale_ub)

    qx_u = qx.to(tl.uint32, bitcast=True)
    sign = qx_u & 0x80000000
    qx_u = qx_u ^ sign
    qx_f = qx_u.to(tl.float32, bitcast=True)

    sat = qx_f >= 6.0
    den = (~sat) & (qx_f < 1.0)
    nor = ~(sat | den)

    den_x = (qx_f + 4194304.0).to(tl.uint32, bitcast=True) - 1249902592
    den_x = den_x.to(tl.uint8)

    mant_odd = (qx_u >> 22) & 1
    nor_x = qx_u + 0xC11FFFFF
    nor_x = nor_x + mant_odd
    nor_x = (nor_x >> 22).to(tl.uint8)

    e2m1 = tl.full([BLOCK_M, NG, SG], 7, dtype=tl.uint8)
    e2m1 = tl.where(nor, nor_x, e2m1)
    e2m1 = tl.where(den, den_x, e2m1)
    e2m1 = e2m1 | (sign >> 28).to(tl.uint8)

    e2m1 = tl.reshape(e2m1, [BLOCK_M, NG, SG // 2, 2])
    ev, od = tl.split(e2m1)
    fp4 = ev | (od << 4)

    return fp4.reshape(BLOCK_M, BLOCK_K // 2), scales.reshape(BLOCK_M, NG)


# =============================================================================
# Hardware MXFP4 quant — used for K > 512, M <= 32 fused path (Path 3)
# =============================================================================

@triton.jit
def _mxfp4_quant_tile_hw(x, BLOCK_M: tl.constexpr, BLOCK_K: tl.constexpr):
    """Quantize (BLOCK_M, BLOCK_K) fp32 tile to MXFP4 via v_cvt_scalef32_pk_fp4_f32.

    The instruction computes fp4(src / scale), so passing 2^scale_ub produces
    the expected block-scaled quantization. The int32 output dtype with a tied
    destination operand matches the packed 32-bit register layout expected by
    the instruction.
    """
    SG: tl.constexpr = 32
    NG: tl.constexpr = BLOCK_K // SG

    x = x.reshape(BLOCK_M, NG, SG)

    amax = tl.max(tl.abs(x), axis=2, keep_dims=True)
    amax_i = amax.to(tl.int32, bitcast=True)
    amax_i = (amax_i + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
    amax = amax_i.to(tl.float32, bitcast=True)
    scale_ub = tl.log2(amax).floor() - 2.0
    scale_ub = tl.clamp(scale_ub, min=-127.0, max=127.0)
    scales = scale_ub.to(tl.uint8) + 127

    hw_scale = tl.exp2(scale_ub)
    hw_scale_broadcast = tl.broadcast_to(hw_scale, (BLOCK_M, NG, SG // 2))

    x_pairs = x.reshape(BLOCK_M, NG, SG // 2, 2)
    x_even, x_odd = tl.split(x_pairs)

    old_vdst = tl.zeros((BLOCK_M, NG, SG // 2), dtype=tl.int32)
    fp4_i32 = tl.inline_asm_elementwise(
        asm="v_cvt_scalef32_pk_fp4_f32 $0, $1, $2, $3",
        constraints="=v,v,v,v,0",
        args=[x_even, x_odd, hw_scale_broadcast, old_vdst],
        dtype=tl.int32,
        is_pure=True,
        pack=1,
    )
    fp4 = (fp4_i32 & 0xFF).to(tl.uint8)

    return fp4.reshape(BLOCK_M, BLOCK_K // 2), scales.reshape(BLOCK_M, NG)


# =============================================================================
# B preshuffle unshuffle helper
# =============================================================================

@triton.jit
def _unshuffle_b_preshuffle(b_wide, BN_GROUPS: tl.constexpr, BLOCK_K: tl.constexpr, BLOCK_N: tl.constexpr):
    """Unshuffle preshuffle-loaded B from (BN_GROUPS, WIDE_K) to (BK//2, BN)."""
    b = b_wide.reshape(BN_GROUPS, BLOCK_K // 64, 2, 16, 16)
    b = b.permute(1, 2, 4, 0, 3)
    return b.reshape(BLOCK_K // 2, BLOCK_N)

@triton.jit
def _load_b_scales_from_preshuffled_generic(
    b_scale_ptr,
    stride_bsn, stride_bsk,
    pid_n,
    scale_k_start,
    BLOCK_N: tl.constexpr,
    BLOCK_K: tl.constexpr,
):
    """Load a preshuffled B-scale tile and unpack it to natural `(BLOCK_N, BLOCK_K // 32)` layout."""
    SG: tl.constexpr = 32
    num_scale_k: tl.constexpr = BLOCK_K // SG

    b_scale_block_n = pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)
    SHUFFLED_SCALE_K: tl.constexpr = num_scale_k * SG
    b_scale_k_offs = tl.arange(0, SHUFFLED_SCALE_K)
    scale_k_start_shuffled = scale_k_start * SG
    b_scale_ptrs = (
        b_scale_ptr
        + b_scale_block_n[:, None] * stride_bsn
        + (scale_k_start_shuffled + b_scale_k_offs[None, :]) * stride_bsk
    )
    return tl.load(b_scale_ptrs).reshape(
        BLOCK_N // 32, BLOCK_K // SG // 8, 4, 16, 2, 2, 1,
    ).permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_N, num_scale_k)


# =============================================================================
# Path 1: Autotuned fused GEMM kernel (K <= 512)
# =============================================================================

_fused_k512_configs = [
    triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32, 'BLOCK_K': 256}, num_warps=4, num_stages=2),
    triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32, 'BLOCK_K': 256}, num_warps=4, num_stages=1),
    triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64, 'BLOCK_K': 256}, num_warps=4, num_stages=2),
    triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64, 'BLOCK_K': 256}, num_warps=4, num_stages=1),
    triton.Config({'BLOCK_M': 32, 'BLOCK_N': 32, 'BLOCK_K': 256}, num_warps=4, num_stages=2),
    triton.Config({'BLOCK_M': 32, 'BLOCK_N': 32, 'BLOCK_K': 256}, num_warps=4, num_stages=1),
    triton.Config({'BLOCK_M': 32, 'BLOCK_N': 64, 'BLOCK_K': 256}, num_warps=4, num_stages=2),
    triton.Config({'BLOCK_M': 32, 'BLOCK_N': 64, 'BLOCK_K': 256}, num_warps=4, num_stages=1),
    triton.Config({'BLOCK_M': 32, 'BLOCK_N': 128, 'BLOCK_K': 256}, num_warps=4, num_stages=2),
    triton.Config({'BLOCK_M': 32, 'BLOCK_N': 128, 'BLOCK_K': 256}, num_warps=8, num_stages=2),
]


@triton.autotune(configs=_fused_k512_configs, key=['M', 'N', 'K'])
@triton.jit
def mxfp4_gemm_fused_k512_kernel(
    a_ptr, b_ptr, c_ptr, b_scale_ptr,
    M, N, K,
    stride_am, stride_ak,
    stride_bk, stride_bn,
    stride_cm, stride_cn,
    stride_bsn, stride_bsk,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
):
    SG: tl.constexpr = 32
    BN_GROUPS: tl.constexpr = BLOCK_N // 16
    WIDE_K: tl.constexpr = BLOCK_K // 2 * 16

    tl.assume(stride_am > 0)
    tl.assume(stride_ak > 0)
    tl.assume(stride_bk > 0)
    tl.assume(stride_bn > 0)
    tl.assume(stride_cm > 0)
    tl.assume(stride_cn > 0)
    tl.assume(stride_bsn > 0)
    tl.assume(stride_bsk > 0)

    pid_mn = tl.program_id(0)
    num_pid_n = tl.cdiv(N, BLOCK_N)
    pid_m = pid_mn // num_pid_n
    pid_n = pid_mn % num_pid_n

    offs_m = (pid_m * BLOCK_M + tl.arange(0, BLOCK_M)) % M

    a_offs_k = tl.arange(0, BLOCK_K)
    a_ptrs = a_ptr + offs_m[:, None] * stride_am + a_offs_k[None, :] * stride_ak

    offs_bn_groups = (pid_n * BN_GROUPS + tl.arange(0, BN_GROUPS)) % (N // 16)
    b_wide_offs = tl.arange(0, WIDE_K)
    b_ptrs = b_ptr + offs_bn_groups[:, None] * stride_bn + b_wide_offs[None, :] * stride_bk

    NUM_SCALE_K: tl.constexpr = BLOCK_K // SG

    accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    num_k_iter = tl.cdiv(K, BLOCK_K)
    scale_k_iter_start = 0
    for _ in range(0, num_k_iter):
        a_bf16 = tl.load(a_ptrs)
        a_fp4, a_scales = _mxfp4_quant_tile(a_bf16.to(tl.float32), BLOCK_M, BLOCK_K)

        b_wide = tl.load(b_ptrs)
        b = _unshuffle_b_preshuffle(b_wide, BN_GROUPS, BLOCK_K, BLOCK_N)

        b_scales = _load_b_scales_from_preshuffled_generic(
            b_scale_ptr, stride_bsn, stride_bsk, pid_n, scale_k_iter_start, BLOCK_N, BLOCK_K,
        )

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

        a_ptrs += BLOCK_K * stride_ak
        b_ptrs += WIDE_K * stride_bk
        scale_k_iter_start += NUM_SCALE_K

    offs_cm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_cn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    c_ptrs = c_ptr + offs_cm[:, None] * stride_cm + offs_cn[None, :] * stride_cn
    c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
    tl.store(c_ptrs, accumulator.to(tl.bfloat16), mask=c_mask)


# =============================================================================
# Path 3: Fused hardware-quant GEMM (K > 512, M <= 32) with split-K
# =============================================================================

_fused_klarge_configs = [
    triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32, 'BLOCK_K': 256, 'SPLIT_K': 8}, num_warps=4, num_stages=2),
    triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64, 'BLOCK_K': 256, 'SPLIT_K': 8}, num_warps=4, num_stages=2),
    triton.Config({'BLOCK_M': 16, 'BLOCK_N': 32, 'BLOCK_K': 256, 'SPLIT_K': 4}, num_warps=4, num_stages=2),
    triton.Config({'BLOCK_M': 16, 'BLOCK_N': 64, 'BLOCK_K': 256, 'SPLIT_K': 4}, num_warps=4, num_stages=2),
    triton.Config({'BLOCK_M': 32, 'BLOCK_N': 32, 'BLOCK_K': 256, 'SPLIT_K': 4}, num_warps=4, num_stages=2),
    triton.Config({'BLOCK_M': 32, 'BLOCK_N': 64, 'BLOCK_K': 256, 'SPLIT_K': 4}, num_warps=4, num_stages=2),
]


@triton.autotune(configs=_fused_klarge_configs, key=['M', 'N', 'K'], reset_to_zero=['c_ptr'])
@triton.jit
def mxfp4_gemm_fused_klarge_kernel(
    a_ptr, b_ptr, c_ptr, b_scale_ptr,
    M, N, K,
    stride_am, stride_ak,
    stride_bk, stride_bn,
    stride_cm, stride_cn,
    stride_bsn, stride_bsk,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
    SPLIT_K: tl.constexpr,
):
    """Fused hardware quant + GEMM for K > 512, M <= 32.

    A single kernel handles quantization and accumulation together, while
    `reset_to_zero` makes autotuned split-K accumulation safe.
    """
    SG: tl.constexpr = 32
    BN_GROUPS: tl.constexpr = BLOCK_N // 16
    WIDE_K: tl.constexpr = BLOCK_K // 2 * 16

    tl.assume(stride_am > 0)
    tl.assume(stride_ak > 0)
    tl.assume(stride_bk > 0)
    tl.assume(stride_bn > 0)
    tl.assume(stride_cm > 0)
    tl.assume(stride_cn > 0)
    tl.assume(stride_bsn > 0)
    tl.assume(stride_bsk > 0)

    pid_mn = tl.program_id(0)
    pid_k = tl.program_id(1)

    num_pid_n = tl.cdiv(N, BLOCK_N)
    pid_m = pid_mn // num_pid_n
    pid_n = pid_mn % num_pid_n

    offs_m = (pid_m * BLOCK_M + tl.arange(0, BLOCK_M)) % M

    k_per_split = tl.cdiv(K, SPLIT_K * BLOCK_K) * BLOCK_K
    k_start = tl.minimum(pid_k * k_per_split, K)
    k_end = tl.minimum(k_start + k_per_split, K)

    a_offs_k = tl.arange(0, BLOCK_K)
    a_ptrs = a_ptr + offs_m[:, None] * stride_am + (k_start + a_offs_k[None, :]) * stride_ak

    offs_bn_groups = (pid_n * BN_GROUPS + tl.arange(0, BN_GROUPS)) % (N // 16)
    b_wide_offs = tl.arange(0, WIDE_K)
    b_ptrs = b_ptr + offs_bn_groups[:, None] * stride_bn + (k_start * 8 + b_wide_offs[None, :]) * stride_bk

    NUM_SCALE_K: tl.constexpr = BLOCK_K // SG

    accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    num_k_iter = tl.cdiv(k_end - k_start, BLOCK_K)
    scale_k_iter_start = k_start // SG
    for _ in range(0, num_k_iter):
        a_bf16 = tl.load(a_ptrs)
        a_fp4, a_scales = _mxfp4_quant_tile_hw(a_bf16.to(tl.float32), BLOCK_M, BLOCK_K)

        b_wide = tl.load(b_ptrs)
        b = _unshuffle_b_preshuffle(b_wide, BN_GROUPS, BLOCK_K, BLOCK_N)

        b_scales = _load_b_scales_from_preshuffled_generic(
            b_scale_ptr, stride_bsn, stride_bsk, pid_n, scale_k_iter_start, BLOCK_N, BLOCK_K,
        )

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

        a_ptrs += BLOCK_K * stride_ak
        b_ptrs += WIDE_K * stride_bk
        scale_k_iter_start += NUM_SCALE_K

    offs_cm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_cn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    c_ptrs = c_ptr + offs_cm[:, None] * stride_cm + offs_cn[None, :] * stride_cn
    c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
    tl.atomic_add(c_ptrs, accumulator, mask=c_mask, sem="relaxed")


# =============================================================================
# Path 3: Very-small-M variant with workspace reduction
# =============================================================================

@triton.jit
def mxfp4_gemm_fused_klarge_workspace_kernel(
    a_ptr, b_ptr, partial_ptr, b_scale_ptr,
    M, N, K,
    stride_am, stride_ak,
    stride_bk, stride_bn,
    stride_ps, stride_pm, stride_pn,
    stride_bsn, stride_bsk,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
    SPLIT_K: tl.constexpr,
):
    """Fused hardware-quant GEMM for very small `M` using a split-K workspace."""
    SG: tl.constexpr = 32
    BN_GROUPS: tl.constexpr = BLOCK_N // 16
    WIDE_K: tl.constexpr = BLOCK_K // 2 * 16

    tl.assume(stride_am > 0)
    tl.assume(stride_ak > 0)
    tl.assume(stride_bk > 0)
    tl.assume(stride_bn > 0)
    tl.assume(stride_ps > 0)
    tl.assume(stride_pm > 0)
    tl.assume(stride_pn > 0)
    tl.assume(stride_bsn > 0)
    tl.assume(stride_bsk > 0)

    pid_mn = tl.program_id(0)
    pid_k = tl.program_id(1)

    num_pid_n = tl.cdiv(N, BLOCK_N)
    pid_m = pid_mn // num_pid_n
    pid_n = pid_mn % num_pid_n

    offs_m = (pid_m * BLOCK_M + tl.arange(0, BLOCK_M)) % M

    k_per_split = tl.cdiv(K, SPLIT_K * BLOCK_K) * BLOCK_K
    k_start = tl.minimum(pid_k * k_per_split, K)
    k_end = tl.minimum(k_start + k_per_split, K)

    a_offs_k = tl.arange(0, BLOCK_K)
    a_ptrs = a_ptr + offs_m[:, None] * stride_am + (k_start + a_offs_k[None, :]) * stride_ak

    offs_bn_groups = (pid_n * BN_GROUPS + tl.arange(0, BN_GROUPS)) % (N // 16)
    b_wide_offs = tl.arange(0, WIDE_K)
    b_ptrs = b_ptr + offs_bn_groups[:, None] * stride_bn + (k_start * 8 + b_wide_offs[None, :]) * stride_bk

    NUM_SCALE_K: tl.constexpr = BLOCK_K // SG

    accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    num_k_iter = tl.cdiv(k_end - k_start, BLOCK_K)
    scale_k_iter_start = k_start // SG
    for _ in range(0, num_k_iter):
        a_bf16 = tl.load(a_ptrs)
        a_fp4, a_scales = _mxfp4_quant_tile_hw(a_bf16.to(tl.float32), BLOCK_M, BLOCK_K)

        b_wide = tl.load(b_ptrs)
        b = _unshuffle_b_preshuffle(b_wide, BN_GROUPS, BLOCK_K, BLOCK_N)

        b_scales = _load_b_scales_from_preshuffled_generic(
            b_scale_ptr, stride_bsn, stride_bsk, pid_n, scale_k_iter_start, BLOCK_N, BLOCK_K,
        )

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

        a_ptrs += BLOCK_K * stride_ak
        b_ptrs += WIDE_K * stride_bk
        scale_k_iter_start += NUM_SCALE_K

    offs_pm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_pn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    partial_ptrs = (
        partial_ptr
        + pid_k * stride_ps
        + offs_pm[:, None] * stride_pm
        + offs_pn[None, :] * stride_pn
    )
    partial_mask = (offs_pm[:, None] < M) & (offs_pn[None, :] < N)
    tl.store(partial_ptrs, accumulator, mask=partial_mask)


@triton.jit
def reduce_splitk_workspace_kernel(
    partial_ptr, c_ptr,
    M, N,
    stride_ps, stride_pm, stride_pn,
    stride_cm, stride_cn,
    SPLIT_K: tl.constexpr,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
):
    """Reduce split-K partials for the very-small-`M` Path 3 workspace."""
    tl.assume(stride_ps > 0)
    tl.assume(stride_pm > 0)
    tl.assume(stride_pn > 0)
    tl.assume(stride_cm > 0)
    tl.assume(stride_cn > 0)

    pid_mn = tl.program_id(0)
    num_pid_n = tl.cdiv(N, BLOCK_N)
    pid_m = pid_mn // num_pid_n
    pid_n = pid_mn % num_pid_n

    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)

    accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
    for split_k_idx in tl.static_range(0, SPLIT_K):
        partial_ptrs = (
            partial_ptr
            + split_k_idx * stride_ps
            + offs_m[:, None] * stride_pm
            + offs_n[None, :] * stride_pn
        )
        accumulator += tl.load(partial_ptrs, mask=mask, other=0.0)

    c_ptrs = c_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
    tl.store(c_ptrs, accumulator.to(tl.bfloat16), mask=mask)


# =============================================================================
# Standalone A quantization kernel (Path 2: K > 512, M > 32)
# =============================================================================

@triton.jit
def _standalone_quant_kernel(
    x_ptr, fp4_ptr, scale_shuffled_ptr,
    M, K,
    stride_xm, stride_xk,
    stride_fm, stride_fk,
    SCALE_N_PAD,
    BLOCK_M: tl.constexpr, BLOCK_K: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_k = tl.program_id(1)

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_k = pid_k * BLOCK_K + tl.arange(0, BLOCK_K)

    x_ptrs = x_ptr + offs_m[:, None] * stride_xm + offs_k[None, :] * stride_xk
    mask = (offs_m[:, None] < M) & (offs_k[None, :] < K)
    x = tl.load(x_ptrs, mask=mask, other=0.0).to(tl.float32)

    fp4, scales = _mxfp4_quant_tile_hw(x, BLOCK_M, BLOCK_K)

    SG: tl.constexpr = 32
    NG: tl.constexpr = BLOCK_K // SG

    fp4_offs = pid_k * (BLOCK_K // 2) + tl.arange(0, BLOCK_K // 2)
    fp4_ptrs = fp4_ptr + offs_m[:, None] * stride_fm + fp4_offs[None, :] * stride_fk
    fp4_mask = (offs_m[:, None] < M) & (fp4_offs[None, :] < K // 2)
    tl.store(fp4_ptrs, fp4, mask=fp4_mask)

    sc_offs = pid_k * NG + tl.arange(0, NG)
    sh_m = offs_m[:, None]
    sh_n = sc_offs[None, :]
    sh_m_block = sh_m // 32
    sh_m_rem = sh_m % 32
    sh_m_hi = sh_m_rem // 16
    sh_m_lo = sh_m_rem % 16
    sh_n_block = sh_n // 8
    sh_n_rem = sh_n % 8
    sh_n_hi = sh_n_rem // 4
    sh_n_lo = sh_n_rem % 4
    sc_ptrs = scale_shuffled_ptr + (
        sh_m_hi
        + sh_n_hi * 2
        + sh_m_lo * 4
        + sh_n_lo * 64
        + sh_n_block * 256
        + sh_m_block * 32 * SCALE_N_PAD
    )
    sc_mask = (offs_m[:, None] < M) & (sc_offs[None, :] < K // SG)
    tl.store(sc_ptrs, scales, mask=sc_mask)


_quant_buffers = {}
_small_m_splitk_buffers = {}
_AITER_ASM_KERNEL_NAME_32X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"


def _choose_large_m_aiter_asm_kernel_name(_N, _K):
    return _AITER_ASM_KERNEL_NAME_32X128

def _run_aiter_large_m_gemm(A, B_shuffle, B_scale_sh):
    """Quantize `A` and run the large-`M` direct AITER FP4 GEMM path."""
    import aiter
    from aiter import dtypes

    A_q_bytes, A_scale_shuffled_bytes = _fast_mxfp4_quant(A)
    A_q = A_q_bytes.view(dtypes.fp4x2)
    A_scale_shuffled = A_scale_shuffled_bytes.view(dtypes.fp8_e8m0)

    M, K = A.shape
    N = B_shuffle.shape[0]
    padded_m = triton.cdiv(M, 32) * 32
    out = torch.empty((padded_m, N), dtype=torch.bfloat16, device=A.device)
    kernel_name = _choose_large_m_aiter_asm_kernel_name(N, K)

    try:
        aiter.gemm_a4w4_asm(
            A_q,
            B_shuffle,
            A_scale_shuffled,
            B_scale_sh,
            out,
            kernel_name,
            None,
            1.0,
            0.0,
            True,
            0,
        )
        return out[:M]
    except Exception:
        return aiter.gemm_a4w4(
            A_q,
            B_shuffle,
            A_scale_shuffled,
            B_scale_sh,
            dtype=dtypes.bf16,
            bpreshuffle=True,
        )


def _fast_mxfp4_quant(A):
    """Standalone MXFP4 quant with pre-allocated buffers."""
    M, K = A.shape
    key = (M, K)
    if key not in _quant_buffers:
        scale_m_pad = triton.cdiv(M, 256) * 256
        scale_n_pad = triton.cdiv(K // 32, 8) * 8
        _quant_buffers[key] = (
            torch.empty((M, K // 2), dtype=torch.uint8, device=A.device),
            torch.full((scale_m_pad, scale_n_pad), 127, dtype=torch.uint8, device=A.device),
        )
    fp4, scale_shuffled = _quant_buffers[key]
    if M < 32:
        BLOCK_M = 16
        BLOCK_K = 32
        num_warps = 1
    else:
        BLOCK_M = 32
        BLOCK_K = 256
        num_warps = 4
    scale_n_pad = scale_shuffled.shape[1]
    grid = (triton.cdiv(M, BLOCK_M), triton.cdiv(K, BLOCK_K))
    _standalone_quant_kernel[grid](
        A, fp4, scale_shuffled, M, K,
        A.stride(0), A.stride(1),
        fp4.stride(0), fp4.stride(1),
        scale_n_pad,
        BLOCK_M=BLOCK_M, BLOCK_K=BLOCK_K,
        num_warps=num_warps,
    )
    return fp4, scale_shuffled


def _run_small_m_workspace_gemm(A, B_sh_wide, B_scale_shuffled):
    """Run the very-small-`M` Path 3 helper that reduces split-K partials explicitly."""
    M, K = A.shape
    N = B_sh_wide.shape[0] * 16
    SPLIT_K = 8
    partial_key = (SPLIT_K, M, N, A.device.index)
    if partial_key not in _small_m_splitk_buffers:
        _small_m_splitk_buffers[partial_key] = torch.empty(
            (SPLIT_K, M, N), dtype=torch.float32, device=A.device
        )
    partial = _small_m_splitk_buffers[partial_key]
    out = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)

    partial_grid = (
        triton.cdiv(M, 16) * triton.cdiv(N, 32),
        SPLIT_K,
    )
    mxfp4_gemm_fused_klarge_workspace_kernel[partial_grid](
        A, B_sh_wide, partial, B_scale_shuffled,
        M, N, K,
        A.stride(0), A.stride(1),
        B_sh_wide.stride(1), B_sh_wide.stride(0),
        partial.stride(0), partial.stride(1), partial.stride(2),
        B_scale_shuffled.stride(0), B_scale_shuffled.stride(1),
        BLOCK_M=16, BLOCK_N=32, BLOCK_K=256, SPLIT_K=SPLIT_K,
        num_warps=4,
        num_stages=2,
    )

    reduce_grid = (
        triton.cdiv(M, 16) * triton.cdiv(N, 128),
    )
    reduce_splitk_workspace_kernel[reduce_grid](
        partial, out,
        M, N,
        partial.stride(0), partial.stride(1), partial.stride(2),
        out.stride(0), out.stride(1),
        SPLIT_K=SPLIT_K,
        BLOCK_M=16, BLOCK_N=128,
        num_warps=4,
        num_stages=2,
    )
    return out


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

    B_scale_raw = B_scale_sh.view(torch.uint8)
    padded_N_scale = B_scale_raw.shape[0]
    padded_K_scale = B_scale_raw.shape[1]
    B_scale_shuffled = B_scale_raw.view(padded_N_scale // 32, padded_K_scale * 32)

    B_sh_bytes = B_shuffle.view(torch.uint8)
    B_sh_wide = B_sh_bytes.reshape(N // 16, (K // 2) * 16)

    if K <= 512:
        # Path 1: fused software quantization plus GEMM.
        C = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
        grid = lambda META: (
            triton.cdiv(M, META['BLOCK_M']) * triton.cdiv(N, META['BLOCK_N']),
        )
        mxfp4_gemm_fused_k512_kernel[grid](
            A, B_sh_wide, C, B_scale_shuffled,
            M, N, K,
            A.stride(0), A.stride(1),
            B_sh_wide.stride(1), B_sh_wide.stride(0),
            C.stride(0), C.stride(1),
            B_scale_shuffled.stride(0), B_scale_shuffled.stride(1),
        )
        return C

    elif M <= 16 and K >= 2048:
        return _run_small_m_workspace_gemm(A, B_sh_wide, B_scale_shuffled)

    elif M <= 32:
        # Path 3: fused hardware quantization plus GEMM with split-K accumulation.
        C = torch.zeros((M, N), dtype=torch.float32, device=A.device)
        grid = lambda META: (
            triton.cdiv(M, META['BLOCK_M']) * triton.cdiv(N, META['BLOCK_N']),
            META['SPLIT_K'],
        )
        mxfp4_gemm_fused_klarge_kernel[grid](
            A, B_sh_wide, C, B_scale_shuffled,
            M, N, K,
            A.stride(0), A.stride(1),
            B_sh_wide.stride(1), B_sh_wide.stride(0),
            C.stride(0), C.stride(1),
            B_scale_shuffled.stride(0), B_scale_shuffled.stride(1),
        )
        return C.to(torch.bfloat16)

    else:
        return _run_aiter_large_m_gemm(A, B_shuffle, B_scale_sh)
scrolls · 693 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 711914.

"""MXFP4 GEMM with shape-specialized dispatch.
- Path 1 (K <= 512): fused Triton kernel that quantizes A in-register and calls
- `tl.dot_scaled`.
+ Path 1 (`K <= 512`) uses a fused Triton kernel that quantizes `A` in-register
+ and calls `tl.dot_scaled`.
- Path 2 (K > 512, M > 32): standalone Triton quantization for A followed by a
- direct AITER FP4 GEMM call on preshuffled B.
+ Path 2 (`K > 512, M > 32`) quantizes `A` with Triton and calls the direct AITER
+ FP4 GEMM on preshuffled `B`.
- Path 3 (K > 512, M <= 32): fused Triton kernel that uses hardware FP4
- conversion for A and split-K accumulation.
+ Path 3 (`K > 512, M <= 32`) keeps quantization and GEMM fused in Triton. The
+ very-small-`M` corner uses a split-K workspace reduction, while the rest of the
+ small-`M` range uses direct split-K accumulation.
"""
import torch
⋯ 119 unchanged lines
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
):
- """Load natural-layout B scales for BK values with contiguous shuffled storage."""
+ """Load a preshuffled B-scale tile and unpack it to natural `(BLOCK_N, BLOCK_K // 32)` layout."""
SG: tl.constexpr = 32
num_scale_k: tl.constexpr = BLOCK_K // SG
⋯ 151 unchanged lines
offs_m = (pid_m * BLOCK_M + tl.arange(0, BLOCK_M)) % M
k_per_split = tl.cdiv(K, SPLIT_K * BLOCK_K) * BLOCK_K
- k_start = pid_k * k_per_split
+ k_start = tl.minimum(pid_k * k_per_split, K)
k_end = tl.minimum(k_start + k_per_split, K)
a_offs_k = tl.arange(0, BLOCK_K)
⋯ 34 unchanged lines
# =============================================================================
+ # Path 3: Very-small-M variant with workspace reduction
+ # =============================================================================
+
+ @triton.jit
+ def mxfp4_gemm_fused_klarge_workspace_kernel(
+ a_ptr, b_ptr, partial_ptr, b_scale_ptr,
+ M, N, K,
+ stride_am, stride_ak,
+ stride_bk, stride_bn,
+ stride_ps, stride_pm, stride_pn,
+ stride_bsn, stride_bsk,
+ BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
+ SPLIT_K: tl.constexpr,
+ ):
+ """Fused hardware-quant GEMM for very small `M` using a split-K workspace."""
+ SG: tl.constexpr = 32
+ BN_GROUPS: tl.constexpr = BLOCK_N // 16
+ WIDE_K: tl.constexpr = BLOCK_K // 2 * 16
+
+ tl.assume(stride_am > 0)
+ tl.assume(stride_ak > 0)
+ tl.assume(stride_bk > 0)
+ tl.assume(stride_bn > 0)
+ tl.assume(stride_ps > 0)
+ tl.assume(stride_pm > 0)
+ tl.assume(stride_pn > 0)
+ tl.assume(stride_bsn > 0)
+ tl.assume(stride_bsk > 0)
+
+ pid_mn = tl.program_id(0)
+ pid_k = tl.program_id(1)
+
+ num_pid_n = tl.cdiv(N, BLOCK_N)
+ pid_m = pid_mn // num_pid_n
+ pid_n = pid_mn % num_pid_n
+
+ offs_m = (pid_m * BLOCK_M + tl.arange(0, BLOCK_M)) % M
+
+ k_per_split = tl.cdiv(K, SPLIT_K * BLOCK_K) * BLOCK_K
+ k_start = tl.minimum(pid_k * k_per_split, K)
+ k_end = tl.minimum(k_start + k_per_split, K)
+
+ a_offs_k = tl.arange(0, BLOCK_K)
+ a_ptrs = a_ptr + offs_m[:, None] * stride_am + (k_start + a_offs_k[None, :]) * stride_ak
+
+ offs_bn_groups = (pid_n * BN_GROUPS + tl.arange(0, BN_GROUPS)) % (N // 16)
+ b_wide_offs = tl.arange(0, WIDE_K)
+ b_ptrs = b_ptr + offs_bn_groups[:, None] * stride_bn + (k_start * 8 + b_wide_offs[None, :]) * stride_bk
+
+ NUM_SCALE_K: tl.constexpr = BLOCK_K // SG
+
+ accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
+
+ num_k_iter = tl.cdiv(k_end - k_start, BLOCK_K)
+ scale_k_iter_start = k_start // SG
+ for _ in range(0, num_k_iter):
+ a_bf16 = tl.load(a_ptrs)
+ a_fp4, a_scales = _mxfp4_quant_tile_hw(a_bf16.to(tl.float32), BLOCK_M, BLOCK_K)
+
+ b_wide = tl.load(b_ptrs)
+ b = _unshuffle_b_preshuffle(b_wide, BN_GROUPS, BLOCK_K, BLOCK_N)
+
+ b_scales = _load_b_scales_from_preshuffled_generic(
+ b_scale_ptr, stride_bsn, stride_bsk, pid_n, scale_k_iter_start, BLOCK_N, BLOCK_K,
+ )
+
+ accumulator += tl.dot_scaled(a_fp4, a_scales, "e2m1", b, b_scales, "e2m1")
+
+ a_ptrs += BLOCK_K * stride_ak
+ b_ptrs += WIDE_K * stride_bk
+ scale_k_iter_start += NUM_SCALE_K
+
+ offs_pm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
+ offs_pn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
+ partial_ptrs = (
+ partial_ptr
+ + pid_k * stride_ps
+ + offs_pm[:, None] * stride_pm
+ + offs_pn[None, :] * stride_pn
+ )
+ partial_mask = (offs_pm[:, None] < M) & (offs_pn[None, :] < N)
+ tl.store(partial_ptrs, accumulator, mask=partial_mask)
+
+
+ @triton.jit
+ def reduce_splitk_workspace_kernel(
+ partial_ptr, c_ptr,
+ M, N,
+ stride_ps, stride_pm, stride_pn,
+ stride_cm, stride_cn,
+ SPLIT_K: tl.constexpr,
+ BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
+ ):
+ """Reduce split-K partials for the very-small-`M` Path 3 workspace."""
+ tl.assume(stride_ps > 0)
+ tl.assume(stride_pm > 0)
+ tl.assume(stride_pn > 0)
+ tl.assume(stride_cm > 0)
+ tl.assume(stride_cn > 0)
+
+ pid_mn = tl.program_id(0)
+ num_pid_n = tl.cdiv(N, BLOCK_N)
+ pid_m = pid_mn // num_pid_n
+ pid_n = pid_mn % num_pid_n
+
+ 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)
+
+ accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
+ for split_k_idx in tl.static_range(0, SPLIT_K):
+ partial_ptrs = (
+ partial_ptr
+ + split_k_idx * stride_ps
+ + offs_m[:, None] * stride_pm
+ + offs_n[None, :] * stride_pn
+ )
+ accumulator += tl.load(partial_ptrs, mask=mask, other=0.0)
+
+ c_ptrs = c_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
+ tl.store(c_ptrs, accumulator.to(tl.bfloat16), mask=mask)
+
+
+ # =============================================================================
# Standalone A quantization kernel (Path 2: K > 512, M > 32)
# =============================================================================
⋯ 50 unchanged lines
_quant_buffers = {}
+ _small_m_splitk_buffers = {}
_AITER_ASM_KERNEL_NAME_32X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
⋯ 1 unchanged lines
return _AITER_ASM_KERNEL_NAME_32X128
def _run_aiter_large_m_gemm(A, B_shuffle, B_scale_sh):
+ """Quantize `A` and run the large-`M` direct AITER FP4 GEMM path."""
import aiter
from aiter import dtypes
⋯ 66 unchanged lines
return fp4, scale_shuffled
+ def _run_small_m_workspace_gemm(A, B_sh_wide, B_scale_shuffled):
+ """Run the very-small-`M` Path 3 helper that reduces split-K partials explicitly."""
+ M, K = A.shape
+ N = B_sh_wide.shape[0] * 16
+ SPLIT_K = 8
+ partial_key = (SPLIT_K, M, N, A.device.index)
+ if partial_key not in _small_m_splitk_buffers:
+ _small_m_splitk_buffers[partial_key] = torch.empty(
+ (SPLIT_K, M, N), dtype=torch.float32, device=A.device
+ )
+ partial = _small_m_splitk_buffers[partial_key]
+ out = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
+
+ partial_grid = (
+ triton.cdiv(M, 16) * triton.cdiv(N, 32),
+ SPLIT_K,
+ )
+ mxfp4_gemm_fused_klarge_workspace_kernel[partial_grid](
+ A, B_sh_wide, partial, B_scale_shuffled,
+ M, N, K,
+ A.stride(0), A.stride(1),
+ B_sh_wide.stride(1), B_sh_wide.stride(0),
+ partial.stride(0), partial.stride(1), partial.stride(2),
+ B_scale_shuffled.stride(0), B_scale_shuffled.stride(1),
+ BLOCK_M=16, BLOCK_N=32, BLOCK_K=256, SPLIT_K=SPLIT_K,
+ num_warps=4,
+ num_stages=2,
+ )
+
+ reduce_grid = (
+ triton.cdiv(M, 16) * triton.cdiv(N, 128),
+ )
+ reduce_splitk_workspace_kernel[reduce_grid](
+ partial, out,
+ M, N,
+ partial.stride(0), partial.stride(1), partial.stride(2),
+ out.stride(0), out.stride(1),
+ SPLIT_K=SPLIT_K,
+ BLOCK_M=16, BLOCK_N=128,
+ num_warps=4,
+ num_stages=2,
+ )
+ return out
+
+
def custom_kernel(data: input_t) -> output_t:
A, _B, B_q, B_shuffle, B_scale_sh = data
M, K = A.shape
⋯ 8 unchanged lines
B_sh_wide = B_sh_bytes.reshape(N // 16, (K // 2) * 16)
if K <= 512:
- # Path 1: fused software quant + GEMM, SK=1
+ # Path 1: fused software quantization plus GEMM.
C = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
grid = lambda META: (
triton.cdiv(M, META['BLOCK_M']) * triton.cdiv(N, META['BLOCK_N']),
⋯ 8 unchanged lines
)
return C
+ elif M <= 16 and K >= 2048:
+ return _run_small_m_workspace_gemm(A, B_sh_wide, B_scale_shuffled)
+
elif M <= 32:
- # Path 3: fused hardware quant + GEMM, SK=4/8, single launch
+ # Path 3: fused hardware quantization plus GEMM with split-K accumulation.
C = torch.zeros((M, N), dtype=torch.float32, device=A.device)
grid = lambda META: (
triton.cdiv(M, META['BLOCK_M']) * triton.cdiv(N, META['BLOCK_N']),
scrolls · 259 diff lines total

Best evidence level for this revision: reported

JSON