Skip to content
KernelIndex
Search⌘K

submission 636064

Divyansh Khanna · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:b6aa5d342367519994484d0d290cccfe25339a6e1261be42ebd5be18269bd5a3
license declaredunknown
license concludedunknown
authorsDivyansh Khanna
imported2026-08-26

Techniques

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

fp4FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.
num-warps = 1NUM_WARPS = 1
split-kThis version (1 kernel launch + optional splitK reduce):
stages = 1NUM_STAGES = 1
tile-m = 64BLOCK_SIZE_M = 64
tile-n = 32BLOCK_SIZE_N = 32

Kernel source

submission.py851 lines
"""
FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.
Quant logic follows aiter op_tests/test_gemm_a4w4.py (get_triton_quant(QuantType.per_1x32)).
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t

# Reuse aiter's quantization math — no point duplicating 100 lines of FP4 bit manipulation
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op


# ─────────────────────────────────────────────────────────────────────────────
# Patched Triton kernel: dynamic_mxfp4_quant with inline e8m0 scale shuffle.
#
# This is aiter's _dynamic_mxfp4_quant_kernel with one addition: the scale
# store section has a SHUFFLE branch that writes scales in the permuted layout
# gemm_a4w4 expects. The shuffle index math is copied verbatim from aiter's
# _fused_rms_mxfp4_quant_kernel (fused_mxfp4_quant.py:173-191).
# ─────────────────────────────────────────────────────────────────────────────

@triton.heuristics(
    {
        "EVEN_M_N": lambda args: args["M"] % args["BLOCK_SIZE_M"] == 0
        and args["N"] % (args["BLOCK_SIZE_N"] * args["NUM_ITER"]) == 0,
    }
)
@triton.jit
def _dynamic_mxfp4_quant_shuffle_kernel(
    # Pointers
    x_ptr,
    x_fp4_ptr,
    bs_ptr,
    # Strides
    stride_x_m_in,
    stride_x_n_in,
    stride_x_fp4_m_in,
    stride_x_fp4_n_in,
    stride_bs_m_in,
    stride_bs_n_in,
    # Problem size
    M,
    N,
    # Compile-time constants
    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,
    SCALING_MODE: tl.constexpr,
    # --- Shuffle support (new vs aiter's original) ---
    SHUFFLE: tl.constexpr,          # Whether to write scales in shuffled order
    SCALE_N_PAD: tl.constexpr,      # Padded scale column count (multiple of 8)
    SCALE_M_PAD: tl.constexpr,      # Padded scale row count (multiple of 256)
):
    pid_m = tl.program_id(0)
    start_n = tl.program_id(1) * NUM_ITER

    # Cast strides to int64 in case M*N > max int32
    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)
    stride_bs_m = tl.cast(stride_bs_m_in, tl.int64)
    stride_bs_n = tl.cast(stride_bs_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):
        # ── Load input tile ──
        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, cache_modifier=".cg").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, cache_modifier=".cg").to(
                tl.float32
            )

        # ── Quantize to MXFP4 (reuses aiter's _mxfp4_quant_op) ──
        out_tensor, bs_e8m0 = _mxfp4_quant_op(
            x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE
        )

        # ── Store quantized FP4 data (unchanged from original) ──
        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)

        # ── Store block scales ──
        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)
        num_bs_cols = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE

        if SHUFFLE:
            # ── Shuffled store ──
            # Copied from aiter's _fused_rms_mxfp4_quant_kernel (lines 173-191).
            # Decomposes the 2D scale index [m, n] into sub-indices matching
            # gemm_a4w4's (16,16) tile access pattern:
            #   m → [m//32, m%32//16, m%16]  (outer block, half-block, inner row)
            #   n → [n//8,  n%8//4,   n%4 ]  (outer block, half-block, inner col)
            # Then interleaves them so that scales for the same GEMM tile are
            # contiguous in memory.
            bs_offs_0 = bs_offs_m[:, None] // 32        # which 32-row block
            bs_offs_1 = bs_offs_m[:, None] % 32
            bs_offs_2 = bs_offs_1 % 16                  # row within 16-row half
            bs_offs_1 = bs_offs_1 // 16                 # which 16-row half (0 or 1)
            bs_offs_3 = bs_offs_n[None, :] // 8          # which 8-col block
            bs_offs_4 = bs_offs_n[None, :] % 8
            bs_offs_5 = bs_offs_4 % 4                    # col within 4-col half
            bs_offs_4 = bs_offs_4 // 4                   # which 4-col half (0 or 1)
            bs_offs = (
                bs_offs_1
                + bs_offs_4 * 2
                + bs_offs_2 * 2 * 2
                + bs_offs_5 * 2 * 2 * 16
                + bs_offs_3 * 2 * 2 * 16 * 4
                + bs_offs_0 * 2 * 16 * SCALE_N_PAD
            )
            # Out-of-bounds scales get value 127 (e8m0 for scale=1.0, i.e. no scaling)
            bs_valid_mask = (bs_offs_m < M)[:, None] & (bs_offs_n < num_bs_cols)[None, :]
            bs_e8m0 = tl.where(bs_valid_mask, bs_e8m0, 127)
            # Store mask covers the full padded region
            bs_mask = (bs_offs_m < SCALE_M_PAD)[:, None] & (
                bs_offs_n < SCALE_N_PAD
            )[None, :]
            tl.store(bs_ptr + bs_offs, bs_e8m0.to(bs_ptr.type.element_ty), mask=bs_mask)
        else:
            # ── Linear store (original behavior) ──
            bs_offs = (
                bs_offs_m[:, None] * stride_bs_m + bs_offs_n[None, :] * stride_bs_n
            )
            if EVEN_M_N:
                tl.store(bs_ptr + bs_offs, bs_e8m0)
            else:
                bs_mask = (bs_offs_m < M)[:, None] & (
                    bs_offs_n < num_bs_cols
                )[None, :]
                tl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask)


def dynamic_mxfp4_quant_shuffled(
    x: torch.Tensor,
    shuffle: bool = True,
    out_fp4: torch.Tensor = None,
    out_scale: torch.Tensor = None,
) -> tuple[torch.Tensor, torch.Tensor]:
    """
    MXFP4 quantization with optional inline scale shuffle.

    Drop-in replacement for:
        x_fp4, bs = dynamic_mxfp4_quant(x)
        if shuffle:
            bs = e8m0_shuffle(bs)

    When shuffle=True, the Triton kernel writes scales directly in the
    permuted layout that gemm_a4w4 expects. This eliminates:
      - The padded tensor allocation in e8m0_shuffle
      - The copy into the padded tensor
      - The permute + .contiguous() (a full read+write of the scale tensor)

    Args:
        x: [M, K] bf16/fp16 input tensor.
        shuffle: If True, write scales in gemm_a4w4's shuffled layout.
        out_fp4: Optional pre-allocated [M, K//2] uint8 output buffer.
        out_scale: Optional pre-allocated scale output buffer.

    Returns:
        (x_fp4, blockscale_e8m0) — same as dynamic_mxfp4_quant, but with
        scales already shuffled when shuffle=True.
    """
    M, N = x.shape
    assert (N // 2) % 2 == 0

    MXFP4_QUANT_BLOCK_SIZE = 32

    if out_fp4 is not None:
        x_fp4 = out_fp4
    else:
        x_fp4 = torch.empty((M, N // 2), dtype=torch.uint8, device=x.device)

    # Scale dimensions
    SCALE_N_valid = triton.cdiv(N, MXFP4_QUANT_BLOCK_SIZE)

    if shuffle:
        # Pad scale dims to multiples required by gemm_a4w4's tile layout:
        #   rows → multiple of 256 (for 32-row blocks × 8 unroll)
        #   cols → multiple of 8   (for 8-col blocks)
        SCALE_M = triton.cdiv(M, 256) * 256
        SCALE_N = triton.cdiv(SCALE_N_valid, 8) * 8
    else:
        SCALE_M = M
        SCALE_N = SCALE_N_valid

    if out_scale is not None:
        blockscale_e8m0 = out_scale
    else:
        blockscale_e8m0 = torch.empty(
            (SCALE_M, SCALE_N), dtype=torch.uint8, device=x.device
        )

    # ── Kernel launch config (same heuristics as aiter's dynamic_mxfp4_quant) ──
    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 = 4
        BLOCK_SIZE_M = 64
        BLOCK_SIZE_N = 64
        NUM_WARPS = 4
        NUM_STAGES = 2
        if N <= 16384:
            BLOCK_SIZE_M = 32
            BLOCK_SIZE_N = 128

    if N <= 1024:
        NUM_ITER = 1
        NUM_STAGES = 1
        NUM_WARPS = 4
        BLOCK_SIZE_N = min(256, triton.next_power_of_2(N))
        BLOCK_SIZE_N = max(32, BLOCK_SIZE_N)
        BLOCK_SIZE_M = min(8, triton.next_power_of_2(M))

    grid = (
        triton.cdiv(M, BLOCK_SIZE_M),
        triton.cdiv(N, BLOCK_SIZE_N * NUM_ITER),
    )

    _dynamic_mxfp4_quant_shuffle_kernel[grid](
        x,
        x_fp4,
        blockscale_e8m0,
        *x.stride(),
        *x_fp4.stride(),
        *blockscale_e8m0.stride(),
        M=M,
        N=N,
        MXFP4_QUANT_BLOCK_SIZE=MXFP4_QUANT_BLOCK_SIZE,
        SCALING_MODE=0,
        NUM_ITER=NUM_ITER,
        BLOCK_SIZE_M=BLOCK_SIZE_M,
        BLOCK_SIZE_N=BLOCK_SIZE_N,
        NUM_STAGES=NUM_STAGES,
        # Shuffle params
        SHUFFLE=shuffle,
        SCALE_N_PAD=SCALE_N,
        SCALE_M_PAD=SCALE_M,
        # Triton runtime config
        num_warps=NUM_WARPS,
        waves_per_eu=0,
        num_stages=1,
    )

    return x_fp4, blockscale_e8m0



# ─────────────────────────────────────────────────────────────────────────────
# Kernel implementations
# ─────────────────────────────────────────────────────────────────────────────

def ref_kernel(data: input_t) -> output_t:
    """
    Reference: MXFP4 per-1x32 quant on A; B_shuffle, B_scale_sh from generate_input.
    gemm_a4w4 with bpreshuffle=True.
    """
    import aiter
    from aiter import QuantType, dtypes
    from aiter.ops.triton.quant import dynamic_mxfp4_quant
    from aiter.utility.fp4_utils import e8m0_shuffle

    def _quant_mxfp4(x, shuffle=True):
        x_fp4, bs_e8m0 = dynamic_mxfp4_quant(x)
        if shuffle:
            bs_e8m0 = e8m0_shuffle(bs_e8m0)
        return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(dtypes.fp8_e8m0)

    A, B, B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous()
    B = B.contiguous()
    m, k = A.shape
    n, _ = B.shape

    A_q, A_scale_sh = _quant_mxfp4(A, shuffle=True)
    out_gemm = aiter.gemm_a4w4(
        A_q,
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )
    return out_gemm


def custom_kernel1(data: input_t) -> output_t:
    """
    Optimized MXFP4 quant + GEMM using aiter's fused_rms_mxfp4_quant.

    Key optimization vs custom_ref_kernel:
    ─────────────────────────────────────
    Reference does 3 steps (3 kernel launches, 2 unnecessary global memory round-trips):
      1. dynamic_mxfp4_quant(A)           → Triton kernel: writes A_q + A_scale to global mem
      2. e8m0_shuffle(A_scale)             → PyTorch ops (alloc padded tensor, copy, view,
                                              permute, .contiguous()) = extra kernel + full
                                              read+write of scale tensor through global memory
      3. gemm_a4w4(...)                    → CK/ASM GEMM: reads A_q + A_scale back

    This version does 2 steps (2 kernel launches):
      1. fused_rms_mxfp4_quant(A, shuffle=True)
                                           → Single Triton kernel that performs:
                                              a) RMSNorm (identity when weight=ones, eps=0,
                                                 and input is already unit-RMS — see note below)
                                              b) MXFP4 quantization
                                              c) e8m0 scale shuffle (inline, no extra alloc)
                                              All in one pass over A, writing shuffled scales
                                              directly without the pad+permute+contiguous dance.
      2. gemm_a4w4(...)                    → CK/ASM GEMM (unchanged)

    IMPORTANT NOTE on correctness:
    ─────────────────────────────
    fused_rms_mxfp4_quant always applies RMSNorm: x_out = x / rms(x) * weight.
    With weight=ones and eps=0, this normalizes each row to unit RMS norm.
    This CHANGES the magnitude of A, so the GEMM output will differ from the
    reference by a per-row scaling factor (rms of each row of A).

    For a truly correct drop-in replacement, we would need either:
      a) A patched dynamic_mxfp4_quant that accepts shuffle=True (the underlying
         Triton kernel _dynamic_mxfp4_quant_kernel does NOT have a SHUFFLE param), or
      b) A standalone Triton wrapper that calls _mxfp4_quant_op + inline shuffle
         without the RMSNorm.

    This implementation demonstrates the fused kernel pattern. If correctness vs
    the reference is required, fall back to custom_ref_kernel until (a) or (b)
    is available.
    """
    import aiter
    from aiter import dtypes
    from aiter.ops.triton.quant.fused_mxfp4_quant import fused_rms_mxfp4_quant

    A, B, B_q, B_shuffle, B_scale_sh = data
    A = A.contiguous()
    m, k = A.shape

    # --- Step 1: Fused quant + shuffle in a single Triton kernel ---
    # fused_rms_mxfp4_quant does: RMSNorm → MXFP4 quant → scale shuffle
    # We pass weight=ones and eps=0 so RMSNorm becomes x / rms(x).
    # The shuffle=True flag writes e8m0 scales in the permuted layout that
    # gemm_a4w4 expects, avoiding the separate e8m0_shuffle() call which
    # would allocate a padded tensor, copy, view(6D), permute, .contiguous().
    ones = torch.ones(k, dtype=A.dtype, device=A.device)
    (A_q, A_scale_sh), _, _, _ = fused_rms_mxfp4_quant(
        A,
        x1_weight=ones,
        x1_epsilon=0.0,
        shuffle=True,          # <-- inline scale shuffle, no separate kernel
    )

    # Reinterpret raw uint8 outputs as the fp4x2 / fp8_e8m0 dtypes
    # that gemm_a4w4 expects (these are zero-cost view operations, no copy).
    A_q = A_q.view(dtypes.fp4x2)
    A_scale_sh = A_scale_sh.view(dtypes.fp8_e8m0)

    # --- Step 2: GEMM (unchanged from reference) ---
    # gemm_a4w4 dispatches to CK (Composable Kernel) or hand-tuned ASM
    # depending on shape and tuning config. bpreshuffle=True tells it that
    # B is already in (16,16)-tile-coalesced layout from shuffle_weight().
    out_gemm = aiter.gemm_a4w4(
        A_q,
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )
    return out_gemm


def custom_kernel2(data: input_t) -> output_t:
    """
    Optimized MXFP4 quant + GEMM — correct drop-in for ref_kernel.

    Uses a patched dynamic_mxfp4_quant that bakes the e8m0 scale shuffle
    into the Triton kernel's store logic. This eliminates the separate
    e8m0_shuffle() call (padded alloc + copy + permute + .contiguous()).

    Reference (3 kernel launches):
      1. dynamic_mxfp4_quant(A)    → Triton: quant, writes A_q + A_scale
      2. e8m0_shuffle(A_scale)     → PyTorch: alloc padded buf, copy, permute, .contiguous()
      3. gemm_a4w4(...)            → CK/ASM GEMM

    This version (2 kernel launches):
      1. dynamic_mxfp4_quant_shuffled(A, shuffle=True)
                                   → Triton: quant + shuffled scale store in one kernel
      2. gemm_a4w4(...)            → CK/ASM GEMM (unchanged)
    """
    import aiter
    from aiter import dtypes

    A, B, B_q, B_shuffle, B_scale_sh = data
    m, k = A.shape

    # Step 1: Quantize A with inline scale shuffle (single Triton kernel).
    # Reuses aiter's _mxfp4_quant_op for the FP4 math, but writes
    # e8m0 block scales directly in gemm_a4w4's permuted layout.
    # No separate e8m0_shuffle needed.
    A_q, A_scale_sh = dynamic_mxfp4_quant_shuffled(A, shuffle=True)

    # Step 2: GEMM (unchanged).
    # .view(dtypes.fp4x2) / .view(dtypes.fp8_e8m0) are zero-cost reinterprets.
    out = aiter.gemm_a4w4(
        A_q.view(dtypes.fp4x2),
        B_shuffle,
        A_scale_sh.view(dtypes.fp8_e8m0),
        B_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )
    return out


def custom_kernel_single_fused(data: input_t) -> output_t:
    """
    Single-kernel fused quant+GEMM using gemm_a16wfp4_preshuffle.

    This eliminates the separate A quantization kernel entirely.
    The preshuffle GEMM kernel quantizes bf16 A to MXFP4 on-the-fly
    in registers during the GEMM loop (via _mxfp4_quant_op + tl.dot_scaled),
    so there are no intermediate FP4 buffers for A written to global memory.

    Reference (3 kernel launches):
      1. dynamic_mxfp4_quant(A)    → writes A_q + A_scale to global mem
      2. e8m0_shuffle(A_scale)     → alloc padded buf, permute, contiguous
      3. gemm_a4w4(...)            → CK/ASM GEMM

    custom_kernel2 (2 kernel launches):
      1. dynamic_mxfp4_quant_shuffled(A)  → quant + shuffled scale store
      2. gemm_a4w4(...)                   → CK/ASM GEMM

    This version (1 kernel launch + optional splitK reduce):
      1. gemm_a16wfp4_preshuffle(A, B, B_scales)
         → Triton GEMM that quantizes A per-tile in registers
         → For small M, uses splitK (parallel K reduction) for better CU utilization

    B format notes:
      - B_shuffle [N, K//2] from shuffle_weight(B_q, (16,16)) must be reshaped
        to [N//16, K//2*16] — same data, different view matching the kernel's
        tiled B loading pattern.
      - B_scale_sh [padded_N, padded_K_scale] from e8m0_shuffle must be reshaped
        to [padded_N//32, padded_K_scale*32] — same shuffled data, different 2D
        view matching the kernel's scale pointer arithmetic. The kernel un-shuffles
        scales in registers (reshape+permute = inverse of e8m0_shuffle).
    """
    A, B, B_q, B_shuffle, B_scale_sh = data
    m, k = A.shape
    n = B.shape[0]

    # Reshape B: [N, K//2] → [N//16, K//2*16]
    B_w = B_shuffle.view(torch.uint8).reshape(n // 16, (k // 2) * 16)

    # Reshape B scales: [padded_N, padded_K_scale] → [padded_N//32, padded_K_scale*32]
    bs = B_scale_sh.view(torch.uint8)
    B_scale_w = bs.reshape(bs.shape[0] // 32, bs.shape[1] * 32)

    configs = {
        # K=512: NUM_KSPLIT=1 (BSK covers full K), no atomic needed
        (2880, 512): {
            "BLOCK_SIZE_M": 8 if m <= 8 else 32,
            "BLOCK_SIZE_N": 64,
            "BLOCK_SIZE_K": 512,
            "GROUP_SIZE_M": 1,
            "num_warps": 4,
            "num_stages": 1,
            "waves_per_eu": 4 if m <= 8 else 2,
            "matrix_instr_nonkdim": 16,
            "cache_modifier": ".cg",
            "NUM_KSPLIT": 1,
        },
        (4096, 512): {
            "BLOCK_SIZE_M": 16 if m <= 32 else 32,
            "BLOCK_SIZE_N": 64,
            "BLOCK_SIZE_K": 512,
            "GROUP_SIZE_M": 1,
            "num_warps": 4,
            "num_stages": 1,
            "waves_per_eu": 4 if m <= 32 else 2,
            "matrix_instr_nonkdim": 16,
            "cache_modifier": ".cg",
            "NUM_KSPLIT": 1,
        },
        # K=7168: splitK=14, atomic eliminates reduce kernel
        (2112, 7168): {
            "BLOCK_SIZE_M": 8 if m <= 8 else (16 if m <= 64 else 32),
            "BLOCK_SIZE_N": 128,
            "BLOCK_SIZE_K": 512,
            "GROUP_SIZE_M": 1,
            "num_warps": 4,
            "num_stages": 1,
            "waves_per_eu": 1 if m <= 8 else 4,
            "matrix_instr_nonkdim": 16,
            "cache_modifier": ".cg",
            "NUM_KSPLIT": 14,
        },
        # K=2048: splitK=4, atomic eliminates reduce kernel
        (7168, 2048): {
            "BLOCK_SIZE_M": 16 if m <= 64 else 32,
            "BLOCK_SIZE_N": 128,
            "BLOCK_SIZE_K": 512,
            "GROUP_SIZE_M": 8 if m >= 128 else 1,
            "num_warps": 4,
            "num_stages": 1,
            "waves_per_eu": 4 if m <= 64 else 2,
            "matrix_instr_nonkdim": 16,
            "cache_modifier": ".cg",
            "NUM_KSPLIT": 4,
        },
        # K=1536: splitK=3
        (3072, 1536): {
            "BLOCK_SIZE_M": 32,
            "BLOCK_SIZE_N": 128,
            "BLOCK_SIZE_K": 512,
            "GROUP_SIZE_M": 8 if m >= 128 else 1,
            "num_warps": 4,
            "num_stages": 1,
            "waves_per_eu": 2,
            "matrix_instr_nonkdim": 16,
            "cache_modifier": ".cg",
            "NUM_KSPLIT": 3,
        },
    }

    config = configs.get((n, k), None)

    # Note: Atomic splitK was attempted (tl.atomic_add to eliminate reduce kernel)
    # but was slower due to torch.zeros init cost, atomic contention, and fp32→bf16
    # conversion overhead. Standard splitK + reduce kernel remains faster.

    from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle

    y = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)

    y = gemm_a16wfp4_preshuffle(
        A, B_w, B_scale_w,
        dtype=torch.bfloat16,
        y=y,
        config=config,
    )

    return y


def custom_kernel(data: input_t) -> output_t:
    """
    Hybrid dispatch: picks the fastest kernel per (m, n, k) shape.

    Three paths:
    - CK/ASM path (gemm_a4w4): hand-tuned assembly, best for some shapes
    - Single-launch Triton (gemm_a16wfp4_preshuffle): fused A quant in registers
    - 2-launch Triton (quant + gemm_afp4wfp4_preshuffle): separate quant, FP4×FP4 GEMM

    The dispatch table below maps each benchmark shape to its fastest path.
    """
    A, B, B_q, B_shuffle, B_scale_sh = data
    m, k = A.shape
    n = B.shape[0]

    # Dispatch key: (m, n, k) for exact match, fallback to (n, k) heuristic
    # Path: "single" = gemm_a16wfp4_preshuffle, "dual" = quant + gemm_afp4wfp4_preshuffle, "ck" = quant + gemm_a4w4
    dispatch = {
        # M=4,  N=2880, K=512:  single-launch (tiny M, fused quant wins)
        (4, 2880, 512): "single",
        # M=16, N=2112, K=7168: single-launch (splitK=14, avoids quant launch)
        (16, 2112, 7168): "single",
        # M=32, N=4096, K=512:  single-launch (small K, 1 tile in K)
        (32, 4096, 512): "single",
        # M=32, N=2880, K=512:  single-launch (small K, 1 tile in K)
        (32, 2880, 512): "single",
        # M=64, N=7168, K=2048: 2-launch Triton (larger M, separate quant is cheaper)
        (64, 7168, 2048): "ck",
        # M=256, N=3072, K=1536: 2-launch Triton (large M, quant kernel efficient)
        (256, 3072, 1536): "ck",
    }

    path = dispatch.get((m, n, k), None)
    if path is None:
        # Heuristic fallback: large M → dual, small M → single
        path = "dual" if m >= 64 else "single"

    if path == "ck":
        # ── CK/ASM path: quant A + gemm_a4w4 (hand-tuned assembly) ──
        import aiter
        from aiter import dtypes

        A_q, A_scale_sh = dynamic_mxfp4_quant_shuffled(A, shuffle=True)
        out = aiter.gemm_a4w4(
            A_q.view(dtypes.fp4x2),
            B_shuffle,
            A_scale_sh.view(dtypes.fp8_e8m0),
            B_scale_sh,
            dtype=dtypes.bf16,
            bpreshuffle=True,
        )
        return out

    # Both Triton paths need reshaped B
    B_w = B_shuffle.view(torch.uint8).reshape(n // 16, (k // 2) * 16)
    bs = B_scale_sh.view(torch.uint8)
    B_scale_w = bs.reshape(bs.shape[0] // 32, bs.shape[1] * 32)

    if path == "dual":
        # ── 2-launch path: quant A separately, then FP4×FP4 GEMM ──
        from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import gemm_afp4wfp4_preshuffle
        from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import _get_config

        configs_2launch = {
            (7168, 2048): {
                "BLOCK_SIZE_M": 32,
                "BLOCK_SIZE_N": 64,
                "BLOCK_SIZE_K": 256,
                "GROUP_SIZE_M": 1,
                "num_warps": 2,
                "num_stages": 2,
                "waves_per_eu": 2,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": ".cg",
                "NUM_KSPLIT": 2,
            },
            (3072, 1536): {
                "BLOCK_SIZE_M": 32,
                "BLOCK_SIZE_N": 128,
                "BLOCK_SIZE_K": 256,
                "GROUP_SIZE_M": 1,
                "num_warps": 4,
                "num_stages": 2,
                "waves_per_eu": 2,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": ".cg",
                "NUM_KSPLIT": 3,
            },
        }

        config = configs_2launch.get((n, k), None)
        if config is None:
            k_internal = k // 2
            config, _ = _get_config(m, n, k_internal, shuffle=True)

        block_m = config["BLOCK_SIZE_M"]

        if block_m >= 32 and m >= 32:
            A_q, A_scale = dynamic_mxfp4_quant_shuffled(A, shuffle=True)
            A_scale_w = A_scale.reshape(A_scale.shape[0] // 32, A_scale.shape[1] * 32)
        else:
            A_q, A_scale = dynamic_mxfp4_quant_shuffled(A, shuffle=False)
            A_scale_w = A_scale

        y = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
        y = gemm_afp4wfp4_preshuffle(
            A_q, B_w, A_scale_w, B_scale_w,
            dtype=torch.bfloat16,
            y=y,
            config=config,
        )
        return y

    else:
        # ── Single-launch path: fused quant+GEMM ──
        from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle

        configs_single = {
            (2880, 512): {
                "BLOCK_SIZE_M": 8 if m <= 8 else 32,
                "BLOCK_SIZE_N": 64,
                "BLOCK_SIZE_K": 512,
                "GROUP_SIZE_M": 1,
                "num_warps": 4,
                "num_stages": 1,
                "waves_per_eu": 4 if m <= 8 else 2,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": ".cg",
                "NUM_KSPLIT": 1,
            },
            (4096, 512): {
                "BLOCK_SIZE_M": 16 if m <= 32 else 32,
                "BLOCK_SIZE_N": 64,
                "BLOCK_SIZE_K": 512,
                "GROUP_SIZE_M": 1,
                "num_warps": 4,
                "num_stages": 1,
                "waves_per_eu": 4 if m <= 32 else 2,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": ".cg",
                "NUM_KSPLIT": 1,
            },
            (2112, 7168): {
                "BLOCK_SIZE_M": 8 if m <= 8 else (16 if m <= 64 else 32),
                "BLOCK_SIZE_N": 128,
                "BLOCK_SIZE_K": 512,
                "GROUP_SIZE_M": 1,
                "num_warps": 4,
                "num_stages": 1,
                "waves_per_eu": 1 if m <= 8 else 4,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": ".cg",
                "NUM_KSPLIT": 14,
            },
            (7168, 2048): {
                "BLOCK_SIZE_M": 16 if m <= 64 else 32,
                "BLOCK_SIZE_N": 128,
                "BLOCK_SIZE_K": 512,
                "GROUP_SIZE_M": 8 if m >= 128 else 1,
                "num_warps": 4,
                "num_stages": 1,
                "waves_per_eu": 4 if m <= 64 else 2,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": ".cg",
                "NUM_KSPLIT": 4,
            },
            (3072, 1536): {
                "BLOCK_SIZE_M": 32,
                "BLOCK_SIZE_N": 128,
                "BLOCK_SIZE_K": 512,
                "GROUP_SIZE_M": 8 if m >= 128 else 1,
                "num_warps": 4,
                "num_stages": 1,
                "waves_per_eu": 2,
                "matrix_instr_nonkdim": 16,
                "cache_modifier": ".cg",
                "NUM_KSPLIT": 3,
            },
        }

        config = configs_single.get((n, k), None)

        y = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
        y = gemm_a16wfp4_preshuffle(
            A, B_w, B_scale_w,
            dtype=torch.bfloat16,
            y=y,
            config=config,
        )
        return y


def custom_kernel3(data: input_t) -> output_t:
    """
    Optimized 2-launch GEMM using gemm_afp4wfp4_preshuffle (Triton FP4×FP4).

    Optimizations vs previous version:
      1. For M >= 32: use shuffle=True in quant kernel to produce preshuffled
         A scales directly, avoiding a separate .permute().contiguous() launch
         (saves 1 kernel launch = ~3-5us)
      2. Custom per-shape configs with tuned NUM_KSPLIT for better CU utilization
      3. Pre-allocated output tensor passed to GEMM to avoid torch.empty overhead
    """
    from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import gemm_afp4wfp4_preshuffle
    from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import _get_config

    A, B, B_q, B_shuffle, B_scale_sh = data
    m, k = A.shape
    n = B.shape[0]

    # Reshape B for preshuffle kernel: [N, K//2] → [N//16, K//2*16]
    B_w = B_shuffle.view(torch.uint8).reshape(n // 16, (k // 2) * 16)

    # Reshape B scales: [padded_N, padded_K_scale] → [padded_N//32, padded_K_scale*32]
    bs = B_scale_sh.view(torch.uint8)
    B_scale_w = bs.reshape(bs.shape[0] // 32, bs.shape[1] * 32)

    # Per-shape configs. Shapes with tuned JSON configs (N=2112,K=7168;
    # N=4096,K=512; N=3072,K=1536) use None → _get_config loads the JSON.
    # Shapes without tuned configs get custom configs here.
    configs = {
        # N=2880, K=512: K_int=256, can't splitK. Use BSN=32 for more tiles.
        # M=4: grid=1*90=90, M=32: grid=1*90=90 (BSM=32,BSN=32)
        (2880, 512): {
            "BLOCK_SIZE_M": 8 if m < 32 else 32,
            "BLOCK_SIZE_N": 32,
            "BLOCK_SIZE_K": 256,
            "GROUP_SIZE_M": 1,
            "num_warps": 2,
            "num_stages": 2,
            "waves_per_eu": 4 if m <= 8 else 2,
            "matrix_instr_nonkdim": 16,
            "cache_modifier": ".cg",
            "NUM_KSPLIT": 1,
        },
        # N=7168, K=2048: K_int=1024. splitK=2 doubles grid tiles.
        # M=64: grid=2*2*112=448 tiles (vs 224 without splitK)
        (7168, 2048): {
            "BLOCK_SIZE_M": 32,
            "BLOCK_SIZE_N": 64,
            "BLOCK_SIZE_K": 256,
            "GROUP_SIZE_M": 1,
            "num_warps": 2,
            "num_stages": 2,
            "waves_per_eu": 2,
            "matrix_instr_nonkdim": 16,
            "cache_modifier": ".cg",
            "NUM_KSPLIT": 2,
        },
    }

    config = configs.get((n, k), None)
    if config is None:
        # Use tuned JSON config (exists for N=2112,K=7168; N=4096,K=512; N=3072,K=1536)
        k_internal = k // 2
        config, _ = _get_config(m, n, k_internal, shuffle=True)

    block_m = config["BLOCK_SIZE_M"]

    # Quantize A with scale format matching what the GEMM kernel expects.
    # Key optimization: for block_m >= 32, we use shuffle=True to produce
    # preshuffled scales directly in the quant kernel, avoiding a separate
    # .permute().contiguous() kernel launch.
    if block_m >= 32 and m >= 32:
        # Shuffle=True: scales written in e8m0_shuffle layout [pad_M, pad_N]
        A_q, A_scale = dynamic_mxfp4_quant_shuffled(A, shuffle=True)
        # Reshape to [pad_M//32, pad_N*32] — zero-cost view, same data layout
        A_scale_w = A_scale.reshape(A_scale.shape[0] // 32, A_scale.shape[1] * 32)
    else:
        # Linear scales for small M (block_m < 32)
        A_q, A_scale = dynamic_mxfp4_quant_shuffled(A, shuffle=False)
        A_scale_w = A_scale

    # Pre-allocate output
    y = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)

    out = gemm_afp4wfp4_preshuffle(
        A_q, B_w, A_scale_w, B_scale_w,
        dtype=torch.bfloat16,
        y=y,
        config=config,
    )
    return out
scrolls · 851 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