Skip to content
KernelIndex
Search⌘K

submission 732442

josusanmartin · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v1156.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-732442?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
7.59µs
#3 of 1143
2026-04-05

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f35b7daf4334f35b89dcf1a4f1f03a4e34dec1c22f2034c110cbd9fea36f756a
license declaredunknown
license concludedunknown
authorsjosusanmartin
imported2026-08-15

Techniques

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

fp4"""Convert 2 FP32 values to packed FP4 byte using hardware instruction.
num-warps = 4num_warps=4, num_stages=1,
split-kfrom aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
stages = 1num_warps=4, num_stages=1,
tile-n = 64ACTUAL_KSPLIT=ACTUAL_KSPLIT, MAX_KSPLIT=MAX_KSPLIT, BLOCK_N=64,

Kernel source

submission_v1156.py919 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""
v1122_cand0873: exact M=16 XCD + M=16 waves=0 + M=256 stages=2 + M=32 stages=3.

Key changes from cand_0902 base:
- exact M=16 XCD kernel (hardcoded KSPLIT=7, NUM_K_ITER=2, pid decomposition)
- M=16 waves_per_eu=0 (auto, was 2)
- M=256 nosplit stages=2 (was 1)
- M=32 shapes stages=3 (was 4)
- duplicate exact M=16 kernel removed (faster JIT)

Measured benchmark result:
- ~7.40 us geomean (non-ranked)
- [5.98, 8.55, 6.16, 6.00, 9.74, 9.81] us
"""
from __future__ import annotations

import os
os.environ["TRITON_HIP_USE_BLOCK_PINGPONG"] = "0"

import torch
import triton
import triton.language as tl
import aiter
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op as _mxfp4_quant_op_sw
from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid, remap_xcd
from task import input_t, output_t


# ---- Hardware FP4 conversion (prescale + scale=1.0 via VGPR) ----

@triton.jit
def _hw_fp4_convert_pair(val0, val1, scale):
    """Convert 2 FP32 values to packed FP4 byte using hardware instruction.
    CK-style: no v_mov_b32, "=v" output constraint, VGPR scale.
    """
    return tl.inline_asm_elementwise(
        "v_cvt_scalef32_pk_fp4_f32 $0, $1, $2, $3",
        "=v,v,v,v",
        [val0, val1, scale],
        dtype=tl.int32,
        is_pure=True,
        pack=1,
    )


@triton.jit
def _mxfp4_quant_op_hw(
    x,
    BLOCK_SIZE_N,
    BLOCK_SIZE_M,
    MXFP4_QUANT_BLOCK_SIZE,
):
    """Hybrid quant: bitwise exponent extraction (no log2/floor) + exp2 prescale.
    Saves 2 SFU ops vs original hw quant while keeping low register pressure."""
    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
    x = x.to(tl.float32).reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)

    # Step 1: Bitwise exponent extraction (avoids log2 + floor)
    amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
    amax_bits = amax.to(tl.int32, bitcast=True)
    amax_bits = (amax_bits + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
    exponent_biased = ((amax_bits >> 23) & 0xFF).to(tl.int32)
    # scale_e8m0_unbiased = exponent_biased - 127 - 2 = exponent_biased - 129
    scale_e8m0_unbiased = exponent_biased - 129
    # tl.clamp doesn't support int32, use tl.where instead
    scale_e8m0_unbiased = tl.where(scale_e8m0_unbiased < -127, -127,
                          tl.where(scale_e8m0_unbiased > 127, 127, scale_e8m0_unbiased))
    bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127

    # Step 2: Prescale via exp2 (keeps low register pressure, avoids bitwise prescale)
    prescale = tl.exp2((-scale_e8m0_unbiased).to(tl.float32))
    x = x * tl.broadcast_to(prescale, x.shape)

    # Step 3: Convert prescaled values to FP4 using hw instruction with scale=1.0
    x_pairs = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2)
    evens, odds = tl.split(x_pairs)
    ones = tl.full(evens.shape, 1.0, dtype=tl.float32)
    packed_i32 = _hw_fp4_convert_pair(evens, odds, ones)
    x_fp4 = (packed_i32 & 0xFF).to(tl.uint8)

    x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
    return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)


@triton.jit
def _mxfp4_quant_op_hw_bitwise(
    x,
    BLOCK_SIZE_N,
    BLOCK_SIZE_M,
    MXFP4_QUANT_BLOCK_SIZE,
):
    """Full bitwise quant — avoids log2, floor, AND exp2 via IEEE 754 manipulation."""
    NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
    x = x.to(tl.float32).reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)

    # Step 1: Compute per-block scale via bitwise exponent extraction
    amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
    amax_bits = amax.to(tl.int32, bitcast=True)
    amax_bits = (amax_bits + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
    # Extract IEEE 754 biased exponent
    exponent_biased = ((amax_bits >> 23) & 0xFF).to(tl.int32)
    # bs_e8m0 = exponent_biased - 2, clamped to [0, 254]
    bs_e8m0_i32 = tl.where(exponent_biased < 2, 0, tl.where(exponent_biased > 256, 254, exponent_biased - 2))
    bs_e8m0 = bs_e8m0_i32.to(tl.uint8)

    # Step 2: Construct prescale as 2^(-scale_e8m0_unbiased) via bit manipulation
    # prescale_exponent_biased = 127 - scale_e8m0_unbiased = 127 - (exponent_biased - 129) = 256 - exponent_biased
    # After clamping: prescale_exponent_biased = 127 - (bs_e8m0_i32 - 127) = 254 - bs_e8m0_i32
    prescale_exp = 254 - bs_e8m0_i32
    prescale_bits = (prescale_exp << 23)
    prescale = prescale_bits.to(tl.float32, bitcast=True)
    x = x * tl.broadcast_to(prescale, x.shape)

    # Step 3: HW convert with scale=1.0
    x_pairs = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2)
    evens, odds = tl.split(x_pairs)
    ones = tl.full(evens.shape, 1.0, dtype=tl.float32)
    packed_i32 = _hw_fp4_convert_pair(evens, odds, ones)
    x_fp4 = (packed_i32 & 0xFF).to(tl.uint8)

    x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
    return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)


# Monkey-patch BITWISE quant for preshuffle kernel path (faster for M=64, M=256)
import aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 as _kernel_module
_kernel_module._mxfp4_quant_op = _mxfp4_quant_op_hw_bitwise

from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import _gemm_afp4wfp4_reduce_kernel
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk

# ---- Enhanced preshuffle kernel: +XCD remap, +.wt store, +accumulator arg ----
@triton.heuristics({
    "EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)
    and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
    and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),
    "GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])
    * triton.cdiv(args["N"], args["BLOCK_SIZE_N"]),
})
@triton.jit
def _enhanced_preshuffle_kernel(
    a_ptr, b_ptr, c_ptr, b_scales_ptr,
    M, N, K,
    stride_am, stride_ak, stride_bn, stride_bk,
    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,
    EVEN_K: tl.constexpr,
    num_warps: tl.constexpr, num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
    GRID_MN: tl.constexpr, PREQUANT: tl.constexpr, cache_modifier: tl.constexpr,
):
    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_bsk > 0)
    tl.assume(stride_bsn > 0)
    pid_unified = tl.program_id(axis=0)
    pid_k = pid_unified % NUM_KSPLIT
    pid = pid_unified // NUM_KSPLIT
    # XCD remap ONLY for small grids ≤90 blocks (K=512 cases)
    # Hurts M=64 (224 blocks) and M=256 (384 blocks)
    # no remap
    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
    SCALE_GROUP_SIZE: tl.constexpr = 32
    if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
        num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
        offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
        offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16
        offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
        a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak)
        offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
        offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
        offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
        b_ptrs = b_ptr + (offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk)
        offs_bsn = (pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, (BLOCK_SIZE_N // 32))) % N
        offs_ks = (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[None, :] * stride_bsk
        accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
        for k in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
            b_scales = (tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
                .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))
            if EVEN_K:
                a_bf16 = tl.load(a_ptrs, cache_modifier=".ca")
                b = tl.load(b_ptrs, cache_modifier=cache_modifier)
            b = (b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
                .permute(0, 1, 4, 2, 3, 5)
                .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
                .trans(1, 0))
            if PREQUANT:
                a, a_scales = _mxfp4_quant_op_hw_bitwise(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
            # CHANGE 2: Pass accumulator as arg (enables HW FMA fusion)
            accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)
            a_ptrs += BLOCK_SIZE_K * stride_ak
            b_ptrs += (BLOCK_SIZE_K // 2) * 16 * 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)
        if GRID_MN <= 224:
            tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")
        else:
            tl.store(c_ptrs, c, mask=c_mask)


@triton.heuristics({
    "EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0),
    "GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])
})
@triton.jit
def _enhanced_preshuffle_exact_m64_kernel(
    a_ptr, b_ptr, c_ptr, b_scales_ptr,
    M, N, K,
    stride_am, stride_ak, stride_bn, stride_bk,
    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,
    EVEN_K: tl.constexpr,
    num_warps: tl.constexpr, num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
    GRID_MN: tl.constexpr, PREQUANT: tl.constexpr, cache_modifier: tl.constexpr,
):
    pid = tl.program_id(axis=0)
    pid_m = pid // 56
    pid_n = pid % 56
    SCALE_GROUP_SIZE: tl.constexpr = 32
    NUM_K_ITER: tl.constexpr = 4
    offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
    offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
    a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_k_bf16[None, :] * stride_ak
    offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
    offs_bn = pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)
    b_ptrs = b_ptr + offs_bn[:, None] * stride_bn + offs_k_shuffle_arr[None, :] * stride_bk
    offs_bsn = pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)
    offs_ks = tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)
    b_scale_ptrs = b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk
    accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
    for _ in range(NUM_K_ITER):
        b_scales = (
            tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
            .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)
        )
        a_bf16 = tl.load(a_ptrs, cache_modifier=".ca")
        b = tl.load(b_ptrs, cache_modifier=cache_modifier)
        b = (
            b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
            .permute(0, 1, 4, 2, 3, 5)
            .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
            .trans(1, 0)
        )
        if PREQUANT:
            a, a_scales = _mxfp4_quant_op_hw_bitwise(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
        accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)
        a_ptrs += BLOCK_SIZE_K * stride_ak
        b_ptrs += (BLOCK_SIZE_K // 2) * 16 * 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, :]
    tl.store(c_ptrs, c, cache_modifier=".wt")


@triton.heuristics({
    "EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0),
    "GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])
})
@triton.jit
def _enhanced_preshuffle_exact_m256_kernel(
    a_ptr, b_ptr, c_ptr, b_scales_ptr,
    M, N, K,
    stride_am, stride_ak, stride_bn, stride_bk,
    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,
    EVEN_K: tl.constexpr,
    num_warps: tl.constexpr, num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
    GRID_MN: tl.constexpr, PREQUANT: tl.constexpr, cache_modifier: tl.constexpr,
):
    pid = tl.program_id(axis=0)
    pid_m = pid // 24
    pid_n = pid % 24
    SCALE_GROUP_SIZE: tl.constexpr = 32
    NUM_K_ITER: tl.constexpr = 6
    offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
    offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
    a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_k_bf16[None, :] * stride_ak
    offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
    offs_bn = pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)
    b_ptrs = b_ptr + offs_bn[:, None] * stride_bn + offs_k_shuffle_arr[None, :] * stride_bk
    offs_bsn = pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)
    offs_ks = tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)
    b_scale_ptrs = b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk
    accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
    for _ in range(NUM_K_ITER):
        b_scales = (
            tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
            .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)
        )
        a_bf16 = tl.load(a_ptrs, cache_modifier=".ca")
        b = tl.load(b_ptrs, cache_modifier=cache_modifier)
        b = (
            b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
            .permute(0, 1, 4, 2, 3, 5)
            .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
            .trans(1, 0)
        )
        if PREQUANT:
            a, a_scales = _mxfp4_quant_op_hw_bitwise(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
        accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)
        a_ptrs += BLOCK_SIZE_K * stride_ak
        b_ptrs += (BLOCK_SIZE_K // 2) * 16 * 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, :]
    tl.store(c_ptrs, c)


# ---- Exact M=16 XCD kernel (hardcoded for M=16, N=2112, KSPLIT=7) ----
@triton.heuristics({"EVEN_K": lambda args: True, "GRID_MN": lambda args: 34})
@triton.jit
def _gemm_exact_m16_xcd_kernel(
    a_ptr, b_ptr, c_ptr, b_scales_ptr,
    M, N, K,
    stride_am, stride_ak, stride_bn, stride_bk,
    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,
    EVEN_K: tl.constexpr,
    num_warps: tl.constexpr, num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
    GRID_MN: tl.constexpr,
    PREQUANT: tl.constexpr, cache_modifier: tl.constexpr,
):
    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_bsk > 0)
    tl.assume(stride_bsn > 0)
    pid_linear = tl.program_id(axis=0)
    pid_k = pid_linear % 7
    pid = remap_xcd(pid_linear // 7, 34, NUM_XCDS=8)
    pid_m = pid % 2
    pid_n = pid // 2
    SCALE_GROUP_SIZE: tl.constexpr = 32
    NUM_K_ITER: tl.constexpr = 2
    offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
    offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16
    offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
    a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak
    offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
    offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
    offs_bn = pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)
    b_ptrs = b_ptr + offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk
    offs_bsn = pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)
    offs_ks = (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[None, :] * stride_bsk
    accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
    for _ in range(NUM_K_ITER):
        b_scales = (
            tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
            .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)
        )
        a_bf16 = tl.load(a_ptrs, cache_modifier=".ca")
        b = tl.load(b_ptrs, cache_modifier=cache_modifier)
        b = (
            b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
            .permute(0, 1, 4, 2, 3, 5)
            .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
            .trans(1, 0)
        )
        a, a_scales = _mxfp4_quant_op_hw(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
        accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)
        a_ptrs += BLOCK_SIZE_K * stride_ak
        b_ptrs += (BLOCK_SIZE_K // 2) * 16 * 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_cn[None, :] < N
    tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")

_BF16 = dtypes.bf16
_FP4X2 = dtypes.fp4x2
_FP8_E8M0 = dtypes.fp8_e8m0
_KERNEL_32X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"

_FUSED_CONFIGS = {
    (4, 2880, 512): {
        "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256,
        "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3,
        "waves_per_eu": 2, "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg", "NUM_KSPLIT": 1,
    },
    (16, 2112, 7168): {
        "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 2, "num_warps": 4, "num_stages": 2,
        "waves_per_eu": 0, "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg", "NUM_KSPLIT": 7,
    },
    (32, 4096, 512): {
        "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256,
        "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3,
        "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg", "NUM_KSPLIT": 1,
    },
    (32, 2880, 512): {
        "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256,
        "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3,
        "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
        "cache_modifier": None, "NUM_KSPLIT": 1,
    },
    # M=256 via fused preshuffle — v925: BM=16 + cache=".cg"
    (256, 3072, 1536): {
        "BLOCK_SIZE_M": 16, "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": 1,
    },
}

_FUSED_K7168_KSPLIT14 = {
    "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
    "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3,
    "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
    "cache_modifier": ".cg", "NUM_KSPLIT": 14,
}

_FUSED_K1536_SPLIT = {
    "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
    "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3,
    "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
    "cache_modifier": ".cg", "NUM_KSPLIT": 3,
}

_FUSED_K1536_NOSPLIT = {
    "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256,
    "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
    "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
    "cache_modifier": ".cg", "NUM_KSPLIT": 1,
}

_FUSED_BM32_K512 = {
    "BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
    "GROUP_SIZE_M": 1, "num_warps": 8, "num_stages": 3,
    "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
    "cache_modifier": ".cg", "NUM_KSPLIT": 1,
}

_FUSED_K2048_NOSPLIT = {
    "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
    "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
    "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
    "cache_modifier": ".cg", "NUM_KSPLIT": 1,
}

_ASM_CONFIGS = {
    # Removed M=256 from ASM — use fused preshuffle instead
}

_QUANT_BLOCK = 32
_QUANT_TILE = 128


# ---- Standalone quant kernel for ASM path (SOFTWARE quant) ----

@triton.jit
def _quant_kernel_asm_layout(
    x_ptr, x_fp4_ptr, bs_ptr,
    stride_x_m, stride_x_n,
    stride_x_fp4_m, stride_x_fp4_n,
    stride_bs_m, stride_bs_n,
    M: tl.constexpr, N: tl.constexpr,
    scaleN: tl.constexpr, scaleM_pad: tl.constexpr, scaleN_pad: tl.constexpr,
    BLOCK_SIZE: tl.constexpr, MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
    SHUFFLE: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)
    stride_x_m = tl.cast(stride_x_m, tl.int64)
    stride_x_n = tl.cast(stride_x_n, tl.int64)
    stride_x_fp4_m = tl.cast(stride_x_fp4_m, tl.int64)
    stride_x_fp4_n = tl.cast(stride_x_fp4_n, tl.int64)
    x_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
    x_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE)
    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).to(tl.float32)
    # Use SOFTWARE quant for ASM path
    out_tensor, bs_e8m0 = _mxfp4_quant_op_sw(
        x, MXFP4_QUANT_BLOCK_SIZE, BLOCK_SIZE, MXFP4_QUANT_BLOCK_SIZE,
    )
    out_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
    out_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE // 2 + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE // 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 + tl.arange(0, BLOCK_SIZE)
    bs_offs_n = pid_n
    if SHUFFLE:
        bs_offs_0 = bs_offs_m[:, None] // 32
        bs_offs_1 = bs_offs_m[:, None] % 32
        bs_offs_2 = bs_offs_1 % 16
        bs_offs_1 = bs_offs_1 // 16
        bs_offs_3 = bs_offs_n[None, :] // 8
        bs_offs_4 = bs_offs_n[None, :] % 8
        bs_offs_5 = bs_offs_4 % 4
        bs_offs_4 = bs_offs_4 // 4
        bs_offs = (
            bs_offs_1 + bs_offs_4 * 2 + bs_offs_2 * 4
            + bs_offs_5 * 64 + bs_offs_3 * 256
            + bs_offs_0 * 32 * scaleN
        )
        bs_mask1 = (bs_offs_m < M)[:, None] & (bs_offs_n < scaleN)[None, :]
        bs_mask2 = (bs_offs_m < scaleM_pad)[:, None] & (bs_offs_n < scaleN_pad)[None, :]
        bs_e8m0 = tl.where(bs_mask1, bs_e8m0, 127)
        tl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask2)
    else:
        bs_offs = bs_offs_m[:, None] * stride_bs_m + bs_offs_n[None, :] * stride_bs_n
        bs_mask = (bs_offs_m < M)[:, None] & (bs_offs_n < N)[None, :]
        tl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask)



# ---- Optimized XCD kernel (fused, uses hw quant) ----

@triton.heuristics(
    {
        "EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)
        and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
        and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),
        "GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])
        * triton.cdiv(args["N"], args["BLOCK_SIZE_N"]),
    }
)
@triton.jit
def _gemm_optimized_xcd_kernel(
    a_ptr, b_ptr, c_ptr, b_scales_ptr,
    M, N, K,
    stride_am, stride_ak, stride_bn, stride_bk,
    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,
    EVEN_K: tl.constexpr,
    num_warps: tl.constexpr, num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
    GRID_MN: tl.constexpr,
    PREQUANT: tl.constexpr, cache_modifier: tl.constexpr,
):
    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_bsk > 0)
    tl.assume(stride_bsn > 0)

    pid_linear = tl.program_id(axis=0)
    pid_k = pid_linear % NUM_KSPLIT
    pid = remap_xcd(pid_linear // NUM_KSPLIT, GRID_MN, NUM_XCDS=8)
    num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
    num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
    pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)

    SCALE_GROUP_SIZE: tl.constexpr = 32

    if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
        num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
        offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
        offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16
        offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
        a_ptrs = a_ptr + (
            offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak
        )
        offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
        offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
        offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
        b_ptrs = b_ptr + (
            offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk
        )
        offs_bsn = (
            pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, (BLOCK_SIZE_N // 32))
        ) % N
        offs_ks = (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[None, :] * stride_bsk
        )
        accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
        for _ in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
            b_scales = (
                tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
                .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)
            )
            if EVEN_K:
                a_bf16 = tl.load(a_ptrs, cache_modifier=".ca")
                b = tl.load(b_ptrs, cache_modifier=cache_modifier)
            b = (
                b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
                .permute(0, 1, 4, 2, 3, 5)
                .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
                .trans(1, 0)
            )
            if PREQUANT:
                a, a_scales = _mxfp4_quant_op_hw(
                    a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32
                )
            accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)
            a_ptrs += BLOCK_SIZE_K * stride_ak
            b_ptrs += (BLOCK_SIZE_K // 2) * 16 * 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)
        if GRID_MN <= 224:
            tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")
        else:
            tl.store(c_ptrs, c, mask=c_mask)


def _e8m0_shuffle_safe(scale: torch.Tensor) -> torch.Tensor:
    m, n = scale.shape
    scale_padded = torch.empty(
        ((m + 255) // 256) * 256, ((n + 7) // 8) * 8,
        dtype=scale.dtype, device=scale.device,
    )
    scale_padded.fill_(0x7F)
    scale_padded[:m, :n] = scale
    sm, sn = scale_padded.shape
    return (
        scale_padded.view(sm // 32, 2, 16, sn // 8, 2, 4)
        .permute(0, 3, 5, 2, 4, 1).contiguous().view(sm, sn)
    )


@triton.jit
def _reduce_m16_splitk7_kernel(
    y_ptr, out_ptr,
    N,
    stride_yk, stride_ym, stride_yn,
    stride_om, stride_on,
    ACTUAL_KSPLIT: tl.constexpr,
    MAX_KSPLIT: tl.constexpr,
    BLOCK_N: tl.constexpr,
):
    pid_n = tl.program_id(axis=0)
    offs_m = tl.arange(0, 16)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    n_mask = offs_n < N
    acc = tl.zeros((16, BLOCK_N), dtype=tl.float32)

    for ks in range(MAX_KSPLIT):
        if ks < ACTUAL_KSPLIT:
            y_ptrs = (
                y_ptr
                + ks * stride_yk
                + offs_m[:, None] * stride_ym
                + offs_n[None, :] * stride_yn
            )
            acc += tl.load(y_ptrs, mask=n_mask[None, :], other=0.0)

    out = acc.to(out_ptr.type.element_ty)
    out_ptrs = out_ptr + offs_m[:, None] * stride_om + offs_n[None, :] * stride_on
    tl.store(out_ptrs, out, mask=n_mask[None, :])


def _safe_wrapper(a, b_shuffle, b_scale_sh):
    a_q_raw, a_scale = dynamic_mxfp4_quant(a.contiguous())
    a_scale_sh = _e8m0_shuffle_safe(a_scale)
    return aiter.gemm_a4w4(
        a_q_raw.view(_FP4X2), b_shuffle,
        a_scale_sh.view(_FP8_E8M0), b_scale_sh,
        dtype=_BF16, bpreshuffle=True,
    )


def _get_route(m, n, k):
    key = (m, n, k)
    if key in _FUSED_CONFIGS:
        return ("fused", _FUSED_CONFIGS[key])
    if key in _ASM_CONFIGS:
        return ("asm", _ASM_CONFIGS[key])
    pair = (n, k)
    if pair == (2112, 7168):
        if m < 16:
            return ("fused", _FUSED_K7168_KSPLIT14)
        return ("fused", _FUSED_CONFIGS[(16, 2112, 7168)])
    if pair == (3072, 1536):
        if m <= 16:
            return ("fused", _FUSED_K1536_SPLIT)
        return ("fused", _FUSED_K1536_NOSPLIT)
    if pair == (2880, 512):
        if m < 16:
            return ("fused", _FUSED_CONFIGS[(4, 2880, 512)])
        if m < 128:
            return ("fused", _FUSED_CONFIGS[(32, 2880, 512)])
        return ("fused", _FUSED_BM32_K512)
    if pair == (4096, 512):
        return ("fused", _FUSED_CONFIGS[(32, 4096, 512)])
    if pair == (7168, 2048):
        return ("fused", _FUSED_K2048_NOSPLIT)
    return None


def _make_asm_handler(m, k, n, device, splitk, kernel_name=_KERNEL_32X128):
    padded_m = ((m + 31) >> 5) << 5
    x_fp4 = torch.empty((padded_m, k >> 1), dtype=torch.uint8, device=device)
    sN = (k + _QUANT_BLOCK - 1) // _QUANT_BLOCK
    sN_pad = ((sN + 7) >> 3) << 3
    sM_pad = ((m + 255) >> 8) << 8
    scale = torch.empty((sM_pad, sN_pad), dtype=torch.uint8, device=device)
    out = torch.empty_strided(
        (padded_m, n), (n + 32, 1), dtype=_BF16, device=device,
    )
    x_fp4_view = x_fp4.view(_FP4X2)
    scale_view = scale.view(_FP8_E8M0)
    q_grid = ((m + _QUANT_TILE - 1) // _QUANT_TILE, sN_pad)
    sfp4_0 = x_fp4.stride(0)
    sfp4_1 = x_fp4.stride(1)
    ssc_0 = scale.stride(0)
    ssc_1 = scale.stride(1)
    out_slice = out[:m]
    quant_launch = _quant_kernel_asm_layout[q_grid]

    def handler(a, b_shuffle, b_scale_sh):
        quant_launch(
            a, x_fp4, scale, k, 1, sfp4_0, sfp4_1, ssc_0, ssc_1,
            M=m, N=k, scaleN=sN,
            scaleM_pad=sM_pad, scaleN_pad=sN_pad,
            BLOCK_SIZE=_QUANT_TILE, MXFP4_QUANT_BLOCK_SIZE=_QUANT_BLOCK, SHUFFLE=True,
        )
        gemm_a4w4_asm(
            x_fp4_view, b_shuffle, scale_view, b_scale_sh,
            out, kernel_name, bpreshuffle=True, log2_k_split=splitk,
        )
        return out_slice

    return handler


def _make_fused_handler(m, n, k, device, raw_config, b_scale_sh):
    config = dict(raw_config)
    K_kernel = k // 2
    if config["NUM_KSPLIT"] > 1:
        SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT = get_splitk(
            K_kernel, config["BLOCK_SIZE_K"], config["NUM_KSPLIT"]
        )
        config["SPLITK_BLOCK_SIZE"] = SPLITK_BLOCK_SIZE
        config["BLOCK_SIZE_K"] = BLOCK_SIZE_K
        config["NUM_KSPLIT"] = NUM_KSPLIT
    if config["BLOCK_SIZE_K"] >= 2 * K_kernel:
        config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * K_kernel)
        config["SPLITK_BLOCK_SIZE"] = 2 * K_kernel
        config["NUM_KSPLIT"] = 1
    config["BLOCK_SIZE_N"] = max(config["BLOCK_SIZE_N"], 32)
    if config["NUM_KSPLIT"] <= 1:
        config["SPLITK_BLOCK_SIZE"] = 2 * K_kernel

    has_splitk = config["NUM_KSPLIT"] > 1
    out = torch.empty((m, n), dtype=_BF16, device=device)
    y_pp = None
    if has_splitk:
        y_pp = torch.empty((config["NUM_KSPLIT"], m, n), dtype=torch.float32, device=device)

    BM = config["BLOCK_SIZE_M"]
    BN = config["BLOCK_SIZE_N"]
    total_tiles = triton.cdiv(m, BM) * triton.cdiv(n, BN)
    grid = (config["NUM_KSPLIT"] * total_tiles,)
    exact_large_route = "m64" if (m, n, k) == (64, 7168, 2048) else ("m256" if (m, n, k) == (256, 3072, 1536) else "")
    stride_ck = 0 if y_pp is None else y_pp.stride(0)
    stride_cm = out.stride(0) if y_pp is None else y_pp.stride(1)
    stride_cn = out.stride(1) if y_pp is None else y_pp.stride(2)
    c_ptr = y_pp if has_splitk else out
    a_s0 = k
    w_s0 = k // 2 * 16
    bss_1 = b_scale_sh.size(1)
    ws_s0 = bss_1 * 32

    use_exact_m16 = has_splitk and (m, n, k) == (16, 2112, 7168)
    if use_exact_m16:
        fused_launch = _gemm_exact_m16_xcd_kernel[grid]
    elif has_splitk:
        fused_launch = _gemm_optimized_xcd_kernel[grid]
    else:
        if exact_large_route == "m64":
            grid = (224,)
            fused_launch = _enhanced_preshuffle_exact_m64_kernel[grid]
        elif exact_large_route == "m256":
            grid = (384,)
            fused_launch = _enhanced_preshuffle_exact_m256_kernel[grid]
        else:
            fused_launch = _enhanced_preshuffle_kernel[grid]

    if has_splitk:
        ACTUAL_KSPLIT = triton.cdiv(K_kernel, config["SPLITK_BLOCK_SIZE"] // 2)
        MAX_KSPLIT = triton.next_power_of_2(config["NUM_KSPLIT"])
        use_specialized_reduce = (m, n, k) == (16, 2112, 7168) and MAX_KSPLIT <= 8
        if use_specialized_reduce:
            reduce_grid = (triton.cdiv(n, 64),)
        else:
            reduce_grid = (triton.cdiv(m, 16), triton.cdiv(n, 64))
            reduce_args = (
                m, n, y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
                out.stride(0), out.stride(1), 16, 64, ACTUAL_KSPLIT, MAX_KSPLIT,
            )
            reduce_launch = _gemm_afp4wfp4_reduce_kernel[reduce_grid]

    if has_splitk:
        def handler(a, b_shuffle, b_scale_sh):
            w = b_shuffle.view(torch.uint8)
            ws = b_scale_sh.view(torch.uint8)
            fused_launch(a, w, c_ptr, ws, m, n, K_kernel, a_s0, 1, w_s0, 1,
                         stride_ck, stride_cm, stride_cn, ws_s0, 1, PREQUANT=True, **config)
            if use_specialized_reduce:
                _reduce_m16_splitk7_kernel[reduce_grid](
                    y_pp, out, n,
                    y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
                    out.stride(0), out.stride(1),
                    ACTUAL_KSPLIT=ACTUAL_KSPLIT, MAX_KSPLIT=MAX_KSPLIT, BLOCK_N=64,
                    num_warps=4, num_stages=1,
                )
            else:
                reduce_launch(y_pp, out, *reduce_args)
            return out
    else:
        def handler(a, b_shuffle, b_scale_sh):
            w = b_shuffle.view(torch.uint8)
            ws = b_scale_sh.view(torch.uint8)
            fused_launch(a, w, c_ptr, ws, m, n, K_kernel, a_s0, 1, w_s0, 1,
                         stride_ck, stride_cm, stride_cn, ws_s0, 1, PREQUANT=True, **config)
            return out

    return handler


_HANDLERS = {}
_last_key = None
_last_handler = None


@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
    global _last_key, _last_handler
    a = data[0]
    m = a.size(0)
    k = a.size(1)
    n = data[1].size(0)
    key = (m, n, k)
    if key == _last_key:
        return _last_handler(a, data[3], data[4])
    b_shuffle = data[3]
    b_scale_sh = data[4]
    if key not in _HANDLERS:
        route = _get_route(m, n, k)
        if route is None:
            _HANDLERS[key] = lambda a, bs, bss: _safe_wrapper(a, bs, bss)
        elif route[0] == "asm":
            splitk, kname = route[1]
            _HANDLERS[key] = _make_asm_handler(m, k, n, a.device, splitk, kname)
        else:
            _HANDLERS[key] = _make_fused_handler(m, n, k, a.device, route[1], b_scale_sh)
    handler = _HANDLERS[key]
    _last_key = key
    _last_handler = handler
    return handler(a, b_shuffle, b_scale_sh)
scrolls · 919 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 534310.

#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
+ """
+ v1122_cand0873: exact M=16 XCD + M=16 waves=0 + M=256 stages=2 + M=32 stages=3.
+
+ Key changes from cand_0902 base:
+ - exact M=16 XCD kernel (hardcoded KSPLIT=7, NUM_K_ITER=2, pid decomposition)
+ - M=16 waves_per_eu=0 (auto, was 2)
+ - M=256 nosplit stages=2 (was 1)
+ - M=32 shapes stages=3 (was 4)
+ - duplicate exact M=16 kernel removed (faster JIT)
+
+ Measured benchmark result:
+ - ~7.40 us geomean (non-ranked)
+ - [5.98, 8.55, 6.16, 6.00, 9.74, 9.81] us
+ """
from __future__ import annotations
- """v423 plus fused routing for the hidden 3072x1536 M=64 family."""
+ import os
+ os.environ["TRITON_HIP_USE_BLOCK_PINGPONG"] = "0"
import torch
import triton
⋯ 2 unchanged lines
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
from aiter.ops.triton.quant import dynamic_mxfp4_quant
- from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op as _mxfp4_quant_op_even
-
+ from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op as _mxfp4_quant_op_sw
+ from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid, remap_xcd
from task import input_t, output_t
+
+ # ---- Hardware FP4 conversion (prescale + scale=1.0 via VGPR) ----
+
+ @triton.jit
+ def _hw_fp4_convert_pair(val0, val1, scale):
+ """Convert 2 FP32 values to packed FP4 byte using hardware instruction.
+ CK-style: no v_mov_b32, "=v" output constraint, VGPR scale.
+ """
+ return tl.inline_asm_elementwise(
+ "v_cvt_scalef32_pk_fp4_f32 $0, $1, $2, $3",
+ "=v,v,v,v",
+ [val0, val1, scale],
+ dtype=tl.int32,
+ is_pure=True,
+ pack=1,
+ )
+
+
+ @triton.jit
+ def _mxfp4_quant_op_hw(
+ x,
+ BLOCK_SIZE_N,
+ BLOCK_SIZE_M,
+ MXFP4_QUANT_BLOCK_SIZE,
+ ):
+ """Hybrid quant: bitwise exponent extraction (no log2/floor) + exp2 prescale.
+ Saves 2 SFU ops vs original hw quant while keeping low register pressure."""
+ NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
+ x = x.to(tl.float32).reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
+
+ # Step 1: Bitwise exponent extraction (avoids log2 + floor)
+ amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
+ amax_bits = amax.to(tl.int32, bitcast=True)
+ amax_bits = (amax_bits + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
+ exponent_biased = ((amax_bits >> 23) & 0xFF).to(tl.int32)
+ # scale_e8m0_unbiased = exponent_biased - 127 - 2 = exponent_biased - 129
+ scale_e8m0_unbiased = exponent_biased - 129
+ # tl.clamp doesn't support int32, use tl.where instead
+ scale_e8m0_unbiased = tl.where(scale_e8m0_unbiased < -127, -127,
+ tl.where(scale_e8m0_unbiased > 127, 127, scale_e8m0_unbiased))
+ bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
+
+ # Step 2: Prescale via exp2 (keeps low register pressure, avoids bitwise prescale)
+ prescale = tl.exp2((-scale_e8m0_unbiased).to(tl.float32))
+ x = x * tl.broadcast_to(prescale, x.shape)
+
+ # Step 3: Convert prescaled values to FP4 using hw instruction with scale=1.0
+ x_pairs = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2)
+ evens, odds = tl.split(x_pairs)
+ ones = tl.full(evens.shape, 1.0, dtype=tl.float32)
+ packed_i32 = _hw_fp4_convert_pair(evens, odds, ones)
+ x_fp4 = (packed_i32 & 0xFF).to(tl.uint8)
+
+ x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
+ return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
+
+
+ @triton.jit
+ def _mxfp4_quant_op_hw_bitwise(
+ x,
+ BLOCK_SIZE_N,
+ BLOCK_SIZE_M,
+ MXFP4_QUANT_BLOCK_SIZE,
+ ):
+ """Full bitwise quant — avoids log2, floor, AND exp2 via IEEE 754 manipulation."""
+ NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
+ x = x.to(tl.float32).reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE)
+
+ # Step 1: Compute per-block scale via bitwise exponent extraction
+ amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
+ amax_bits = amax.to(tl.int32, bitcast=True)
+ amax_bits = (amax_bits + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
+ # Extract IEEE 754 biased exponent
+ exponent_biased = ((amax_bits >> 23) & 0xFF).to(tl.int32)
+ # bs_e8m0 = exponent_biased - 2, clamped to [0, 254]
+ bs_e8m0_i32 = tl.where(exponent_biased < 2, 0, tl.where(exponent_biased > 256, 254, exponent_biased - 2))
+ bs_e8m0 = bs_e8m0_i32.to(tl.uint8)
+
+ # Step 2: Construct prescale as 2^(-scale_e8m0_unbiased) via bit manipulation
+ # prescale_exponent_biased = 127 - scale_e8m0_unbiased = 127 - (exponent_biased - 129) = 256 - exponent_biased
+ # After clamping: prescale_exponent_biased = 127 - (bs_e8m0_i32 - 127) = 254 - bs_e8m0_i32
+ prescale_exp = 254 - bs_e8m0_i32
+ prescale_bits = (prescale_exp << 23)
+ prescale = prescale_bits.to(tl.float32, bitcast=True)
+ x = x * tl.broadcast_to(prescale, x.shape)
+
+ # Step 3: HW convert with scale=1.0
+ x_pairs = x.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS, MXFP4_QUANT_BLOCK_SIZE // 2, 2)
+ evens, odds = tl.split(x_pairs)
+ ones = tl.full(evens.shape, 1.0, dtype=tl.float32)
+ packed_i32 = _hw_fp4_convert_pair(evens, odds, ones)
+ x_fp4 = (packed_i32 & 0xFF).to(tl.uint8)
+
+ x_fp4 = x_fp4.reshape(BLOCK_SIZE_M, BLOCK_SIZE_N // 2)
+ return x_fp4, bs_e8m0.reshape(BLOCK_SIZE_M, NUM_QUANT_BLOCKS)
+
+
+ # Monkey-patch BITWISE quant for preshuffle kernel path (faster for M=64, M=256)
import aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 as _kernel_module
+ _kernel_module._mxfp4_quant_op = _mxfp4_quant_op_hw_bitwise
- _kernel_module._mxfp4_quant_op = _mxfp4_quant_op_even
+ from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import _gemm_afp4wfp4_reduce_kernel
+ from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
- from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
+ # ---- Enhanced preshuffle kernel: +XCD remap, +.wt store, +accumulator arg ----
+ @triton.heuristics({
+ "EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)
+ and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
+ and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),
+ "GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])
+ * triton.cdiv(args["N"], args["BLOCK_SIZE_N"]),
+ })
+ @triton.jit
+ def _enhanced_preshuffle_kernel(
+ a_ptr, b_ptr, c_ptr, b_scales_ptr,
+ M, N, K,
+ stride_am, stride_ak, stride_bn, stride_bk,
+ 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,
+ EVEN_K: tl.constexpr,
+ num_warps: tl.constexpr, num_stages: tl.constexpr,
+ waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
+ GRID_MN: tl.constexpr, PREQUANT: tl.constexpr, cache_modifier: tl.constexpr,
+ ):
+ 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_bsk > 0)
+ tl.assume(stride_bsn > 0)
+ pid_unified = tl.program_id(axis=0)
+ pid_k = pid_unified % NUM_KSPLIT
+ pid = pid_unified // NUM_KSPLIT
+ # XCD remap ONLY for small grids ≤90 blocks (K=512 cases)
+ # Hurts M=64 (224 blocks) and M=256 (384 blocks)
+ # no remap
+ 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
+ SCALE_GROUP_SIZE: tl.constexpr = 32
+ if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
+ num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
+ offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
+ offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16
+ offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
+ a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak)
+ offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
+ offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
+ offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
+ b_ptrs = b_ptr + (offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk)
+ offs_bsn = (pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, (BLOCK_SIZE_N // 32))) % N
+ offs_ks = (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[None, :] * stride_bsk
+ accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
+ for k in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
+ b_scales = (tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
+ .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))
+ if EVEN_K:
+ a_bf16 = tl.load(a_ptrs, cache_modifier=".ca")
+ b = tl.load(b_ptrs, cache_modifier=cache_modifier)
+ b = (b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
+ .permute(0, 1, 4, 2, 3, 5)
+ .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
+ .trans(1, 0))
+ if PREQUANT:
+ a, a_scales = _mxfp4_quant_op_hw_bitwise(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
+ # CHANGE 2: Pass accumulator as arg (enables HW FMA fusion)
+ accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)
+ a_ptrs += BLOCK_SIZE_K * stride_ak
+ b_ptrs += (BLOCK_SIZE_K // 2) * 16 * 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)
+ if GRID_MN <= 224:
+ tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")
+ else:
+ tl.store(c_ptrs, c, mask=c_mask)
+
+ @triton.heuristics({
+ "EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0),
+ "GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])
+ })
+ @triton.jit
+ def _enhanced_preshuffle_exact_m64_kernel(
+ a_ptr, b_ptr, c_ptr, b_scales_ptr,
+ M, N, K,
+ stride_am, stride_ak, stride_bn, stride_bk,
+ 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,
+ EVEN_K: tl.constexpr,
+ num_warps: tl.constexpr, num_stages: tl.constexpr,
+ waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
+ GRID_MN: tl.constexpr, PREQUANT: tl.constexpr, cache_modifier: tl.constexpr,
+ ):
+ pid = tl.program_id(axis=0)
+ pid_m = pid // 56
+ pid_n = pid % 56
+ SCALE_GROUP_SIZE: tl.constexpr = 32
+ NUM_K_ITER: tl.constexpr = 4
+ offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
+ offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
+ a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_k_bf16[None, :] * stride_ak
+ offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
+ offs_bn = pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)
+ b_ptrs = b_ptr + offs_bn[:, None] * stride_bn + offs_k_shuffle_arr[None, :] * stride_bk
+ offs_bsn = pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)
+ offs_ks = tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)
+ b_scale_ptrs = b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk
+ accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
+ for _ in range(NUM_K_ITER):
+ b_scales = (
+ tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
+ .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)
+ )
+ a_bf16 = tl.load(a_ptrs, cache_modifier=".ca")
+ b = tl.load(b_ptrs, cache_modifier=cache_modifier)
+ b = (
+ b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
+ .permute(0, 1, 4, 2, 3, 5)
+ .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
+ .trans(1, 0)
+ )
+ if PREQUANT:
+ a, a_scales = _mxfp4_quant_op_hw_bitwise(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
+ accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)
+ a_ptrs += BLOCK_SIZE_K * stride_ak
+ b_ptrs += (BLOCK_SIZE_K // 2) * 16 * 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, :]
+ tl.store(c_ptrs, c, cache_modifier=".wt")
+
+
+ @triton.heuristics({
+ "EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0),
+ "GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])
+ })
+ @triton.jit
+ def _enhanced_preshuffle_exact_m256_kernel(
+ a_ptr, b_ptr, c_ptr, b_scales_ptr,
+ M, N, K,
+ stride_am, stride_ak, stride_bn, stride_bk,
+ 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,
+ EVEN_K: tl.constexpr,
+ num_warps: tl.constexpr, num_stages: tl.constexpr,
+ waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
+ GRID_MN: tl.constexpr, PREQUANT: tl.constexpr, cache_modifier: tl.constexpr,
+ ):
+ pid = tl.program_id(axis=0)
+ pid_m = pid // 24
+ pid_n = pid % 24
+ SCALE_GROUP_SIZE: tl.constexpr = 32
+ NUM_K_ITER: tl.constexpr = 6
+ offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
+ offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
+ a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_k_bf16[None, :] * stride_ak
+ offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
+ offs_bn = pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)
+ b_ptrs = b_ptr + offs_bn[:, None] * stride_bn + offs_k_shuffle_arr[None, :] * stride_bk
+ offs_bsn = pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)
+ offs_ks = tl.arange(0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32)
+ b_scale_ptrs = b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk
+ accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
+ for _ in range(NUM_K_ITER):
+ b_scales = (
+ tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
+ .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)
+ )
+ a_bf16 = tl.load(a_ptrs, cache_modifier=".ca")
+ b = tl.load(b_ptrs, cache_modifier=cache_modifier)
+ b = (
+ b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
+ .permute(0, 1, 4, 2, 3, 5)
+ .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
+ .trans(1, 0)
+ )
+ if PREQUANT:
+ a, a_scales = _mxfp4_quant_op_hw_bitwise(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
+ accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)
+ a_ptrs += BLOCK_SIZE_K * stride_ak
+ b_ptrs += (BLOCK_SIZE_K // 2) * 16 * 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, :]
+ tl.store(c_ptrs, c)
+
+
+ # ---- Exact M=16 XCD kernel (hardcoded for M=16, N=2112, KSPLIT=7) ----
+ @triton.heuristics({"EVEN_K": lambda args: True, "GRID_MN": lambda args: 34})
+ @triton.jit
+ def _gemm_exact_m16_xcd_kernel(
+ a_ptr, b_ptr, c_ptr, b_scales_ptr,
+ M, N, K,
+ stride_am, stride_ak, stride_bn, stride_bk,
+ 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,
+ EVEN_K: tl.constexpr,
+ num_warps: tl.constexpr, num_stages: tl.constexpr,
+ waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
+ GRID_MN: tl.constexpr,
+ PREQUANT: tl.constexpr, cache_modifier: tl.constexpr,
+ ):
+ 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_bsk > 0)
+ tl.assume(stride_bsn > 0)
+ pid_linear = tl.program_id(axis=0)
+ pid_k = pid_linear % 7
+ pid = remap_xcd(pid_linear // 7, 34, NUM_XCDS=8)
+ pid_m = pid % 2
+ pid_n = pid // 2
+ SCALE_GROUP_SIZE: tl.constexpr = 32
+ NUM_K_ITER: tl.constexpr = 2
+ offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
+ offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16
+ offs_am = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
+ a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak
+ offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
+ offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
+ offs_bn = pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)
+ b_ptrs = b_ptr + offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk
+ offs_bsn = pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, BLOCK_SIZE_N // 32)
+ offs_ks = (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[None, :] * stride_bsk
+ accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
+ for _ in range(NUM_K_ITER):
+ b_scales = (
+ tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
+ .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)
+ )
+ a_bf16 = tl.load(a_ptrs, cache_modifier=".ca")
+ b = tl.load(b_ptrs, cache_modifier=cache_modifier)
+ b = (
+ b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
+ .permute(0, 1, 4, 2, 3, 5)
+ .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
+ .trans(1, 0)
+ )
+ a, a_scales = _mxfp4_quant_op_hw(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
+ accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)
+ a_ptrs += BLOCK_SIZE_K * stride_ak
+ b_ptrs += (BLOCK_SIZE_K // 2) * 16 * 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_cn[None, :] < N
+ tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")
+
_BF16 = dtypes.bf16
_FP4X2 = dtypes.fp4x2
_FP8_E8M0 = dtypes.fp8_e8m0
_KERNEL_32X128 = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
- _PUBLIC_SMALL = {
+ _FUSED_CONFIGS = {
(4, 2880, 512): {
- "BLOCK_SIZE_M": 8,
- "BLOCK_SIZE_N": 64,
- "BLOCK_SIZE_K": 512,
- "GROUP_SIZE_M": 1,
- "num_warps": 4,
- "num_stages": 2,
- "waves_per_eu": 2,
- "matrix_instr_nonkdim": 16,
- "cache_modifier": ".cg",
- "NUM_KSPLIT": 1,
+ "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256,
+ "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3,
+ "waves_per_eu": 2, "matrix_instr_nonkdim": 16,
+ "cache_modifier": ".cg", "NUM_KSPLIT": 1,
},
(16, 2112, 7168): {
- "BLOCK_SIZE_M": 16,
- "BLOCK_SIZE_N": 64,
- "BLOCK_SIZE_K": 512,
- "GROUP_SIZE_M": 1,
- "num_warps": 4,
- "num_stages": 2,
- "waves_per_eu": 1,
- "matrix_instr_nonkdim": 16,
- "cache_modifier": ".cg",
- "NUM_KSPLIT": 8,
+ "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
+ "GROUP_SIZE_M": 2, "num_warps": 4, "num_stages": 2,
+ "waves_per_eu": 0, "matrix_instr_nonkdim": 16,
+ "cache_modifier": ".cg", "NUM_KSPLIT": 7,
},
(32, 4096, 512): {
- "BLOCK_SIZE_M": 8,
- "BLOCK_SIZE_N": 64,
- "BLOCK_SIZE_K": 512,
- "GROUP_SIZE_M": 1,
- "num_warps": 4,
- "num_stages": 3,
- "waves_per_eu": 1,
- "matrix_instr_nonkdim": 16,
- "cache_modifier": ".cg",
- "NUM_KSPLIT": 8,
+ "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256,
+ "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3,
+ "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
+ "cache_modifier": ".cg", "NUM_KSPLIT": 1,
},
(32, 2880, 512): {
- "BLOCK_SIZE_M": 8,
- "BLOCK_SIZE_N": 64,
- "BLOCK_SIZE_K": 512,
- "GROUP_SIZE_M": 1,
- "num_warps": 4,
- "num_stages": 2,
- "waves_per_eu": 1,
- "matrix_instr_nonkdim": 16,
- "cache_modifier": None,
- "NUM_KSPLIT": 1,
+ "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 256,
+ "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3,
+ "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
+ "cache_modifier": None, "NUM_KSPLIT": 1,
},
+ # M=256 via fused preshuffle — v925: BM=16 + cache=".cg"
+ (256, 3072, 1536): {
+ "BLOCK_SIZE_M": 16, "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": 1,
+ },
}
- _PUBLIC_LARGE = {
- (64, 7168, 2048): 2,
- (256, 3072, 1536): 1,
+ _FUSED_K7168_KSPLIT14 = {
+ "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
+ "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3,
+ "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
+ "cache_modifier": ".cg", "NUM_KSPLIT": 14,
}
- _PUBLIC_TEST_SMALL = {
- (8, 2112, 7168): {
- "BLOCK_SIZE_M": 8,
- "BLOCK_SIZE_N": 128,
- "BLOCK_SIZE_K": 512,
- "GROUP_SIZE_M": 1,
- "num_warps": 4,
- "num_stages": 1,
- "waves_per_eu": 1,
- "matrix_instr_nonkdim": 16,
- "cache_modifier": ".cg",
- "NUM_KSPLIT": 14,
- },
- (16, 3072, 1536): {
- "BLOCK_SIZE_M": 16,
- "BLOCK_SIZE_N": 128,
- "BLOCK_SIZE_K": 512,
- "GROUP_SIZE_M": 1,
- "num_warps": 4,
- "num_stages": 1,
- "waves_per_eu": 1,
- "matrix_instr_nonkdim": 16,
- "cache_modifier": ".cg",
- "NUM_KSPLIT": 3,
- },
+ _FUSED_K1536_SPLIT = {
+ "BLOCK_SIZE_M": 8, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
+ "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 3,
+ "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
+ "cache_modifier": ".cg", "NUM_KSPLIT": 3,
}
+ _FUSED_K1536_NOSPLIT = {
+ "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 256,
+ "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
+ "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
+ "cache_modifier": ".cg", "NUM_KSPLIT": 1,
+ }
+
+ _FUSED_BM32_K512 = {
+ "BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
+ "GROUP_SIZE_M": 1, "num_warps": 8, "num_stages": 3,
+ "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
+ "cache_modifier": ".cg", "NUM_KSPLIT": 1,
+ }
+
+ _FUSED_K2048_NOSPLIT = {
+ "BLOCK_SIZE_M": 16, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 512,
+ "GROUP_SIZE_M": 1, "num_warps": 4, "num_stages": 2,
+ "waves_per_eu": 1, "matrix_instr_nonkdim": 16,
+ "cache_modifier": ".cg", "NUM_KSPLIT": 1,
+ }
+
+ _ASM_CONFIGS = {
+ # Removed M=256 from ASM — use fused preshuffle instead
+ }
+
_QUANT_BLOCK = 32
_QUANT_TILE = 128
- _OUT_PAD_BF16 = 32
- _BUFS = {}
+ # ---- Standalone quant kernel for ASM path (SOFTWARE quant) ----
+
@triton.jit
- def _dynamic_mxfp4_quant_kernel_even_asm_layout(
- x_ptr,
- x_fp4_ptr,
- bs_ptr,
- stride_x_m,
- stride_x_n,
- stride_x_fp4_m,
- stride_x_fp4_n,
- stride_bs_m,
- stride_bs_n,
- M: tl.constexpr,
- N: tl.constexpr,
- scaleN: tl.constexpr,
- scaleM_pad: tl.constexpr,
- scaleN_pad: tl.constexpr,
- BLOCK_SIZE: tl.constexpr,
- MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
+ def _quant_kernel_asm_layout(
+ x_ptr, x_fp4_ptr, bs_ptr,
+ stride_x_m, stride_x_n,
+ stride_x_fp4_m, stride_x_fp4_n,
+ stride_bs_m, stride_bs_n,
+ M: tl.constexpr, N: tl.constexpr,
+ scaleN: tl.constexpr, scaleM_pad: tl.constexpr, scaleN_pad: tl.constexpr,
+ BLOCK_SIZE: tl.constexpr, MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
SHUFFLE: tl.constexpr,
):
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
-
stride_x_m = tl.cast(stride_x_m, tl.int64)
stride_x_n = tl.cast(stride_x_n, tl.int64)
stride_x_fp4_m = tl.cast(stride_x_fp4_m, tl.int64)
stride_x_fp4_n = tl.cast(stride_x_fp4_n, tl.int64)
-
x_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
x_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE)
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).to(tl.float32)
-
- out_tensor, bs_e8m0 = _mxfp4_quant_op_even(
- x,
- MXFP4_QUANT_BLOCK_SIZE,
- BLOCK_SIZE,
- MXFP4_QUANT_BLOCK_SIZE,
+ # Use SOFTWARE quant for ASM path
+ out_tensor, bs_e8m0 = _mxfp4_quant_op_sw(
+ x, MXFP4_QUANT_BLOCK_SIZE, BLOCK_SIZE, MXFP4_QUANT_BLOCK_SIZE,
)
-
out_offs_m = pid_m * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
- out_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE // 2 + tl.arange(
- 0, MXFP4_QUANT_BLOCK_SIZE // 2
- )
- out_offs = (
- out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
- )
+ out_offs_n = pid_n * MXFP4_QUANT_BLOCK_SIZE // 2 + tl.arange(0, MXFP4_QUANT_BLOCK_SIZE // 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 + tl.arange(0, BLOCK_SIZE)
bs_offs_n = pid_n
-
if SHUFFLE:
bs_offs_0 = bs_offs_m[:, None] // 32
bs_offs_1 = bs_offs_m[:, None] % 32
⋯ 4 unchanged lines
bs_offs_5 = bs_offs_4 % 4
bs_offs_4 = bs_offs_4 // 4
bs_offs = (
- bs_offs_1
- + bs_offs_4 * 2
- + bs_offs_2 * 4
- + bs_offs_5 * 64
- + bs_offs_3 * 256
+ bs_offs_1 + bs_offs_4 * 2 + bs_offs_2 * 4
+ + bs_offs_5 * 64 + bs_offs_3 * 256
+ bs_offs_0 * 32 * scaleN
)
bs_mask1 = (bs_offs_m < M)[:, None] & (bs_offs_n < scaleN)[None, :]
⋯ 6 unchanged lines
tl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask)
+
+ # ---- Optimized XCD kernel (fused, uses hw quant) ----
+
+ @triton.heuristics(
+ {
+ "EVEN_K": lambda args: (args["K"] % (args["BLOCK_SIZE_K"] // 2) == 0)
+ and (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
+ and (args["K"] % (args["SPLITK_BLOCK_SIZE"] // 2) == 0),
+ "GRID_MN": lambda args: triton.cdiv(args["M"], args["BLOCK_SIZE_M"])
+ * triton.cdiv(args["N"], args["BLOCK_SIZE_N"]),
+ }
+ )
+ @triton.jit
+ def _gemm_optimized_xcd_kernel(
+ a_ptr, b_ptr, c_ptr, b_scales_ptr,
+ M, N, K,
+ stride_am, stride_ak, stride_bn, stride_bk,
+ 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,
+ EVEN_K: tl.constexpr,
+ num_warps: tl.constexpr, num_stages: tl.constexpr,
+ waves_per_eu: tl.constexpr, matrix_instr_nonkdim: tl.constexpr,
+ GRID_MN: tl.constexpr,
+ PREQUANT: tl.constexpr, cache_modifier: tl.constexpr,
+ ):
+ 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_bsk > 0)
+ tl.assume(stride_bsn > 0)
+
+ pid_linear = tl.program_id(axis=0)
+ pid_k = pid_linear % NUM_KSPLIT
+ pid = remap_xcd(pid_linear // NUM_KSPLIT, GRID_MN, NUM_XCDS=8)
+ num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
+ num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
+ pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
+
+ SCALE_GROUP_SIZE: tl.constexpr = 32
+
+ if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
+ num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
+ offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
+ offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16
+ offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
+ a_ptrs = a_ptr + (
+ offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak
+ )
+ offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
+ offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
+ offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
+ b_ptrs = b_ptr + (
+ offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk
+ )
+ offs_bsn = (
+ pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, (BLOCK_SIZE_N // 32))
+ ) % N
+ offs_ks = (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[None, :] * stride_bsk
+ )
+ accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
+ for _ in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
+ b_scales = (
+ tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
+ .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)
+ )
+ if EVEN_K:
+ a_bf16 = tl.load(a_ptrs, cache_modifier=".ca")
+ b = tl.load(b_ptrs, cache_modifier=cache_modifier)
+ b = (
+ b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
+ .permute(0, 1, 4, 2, 3, 5)
+ .reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
+ .trans(1, 0)
+ )
+ if PREQUANT:
+ a, a_scales = _mxfp4_quant_op_hw(
+ a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32
+ )
+ accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)
+ a_ptrs += BLOCK_SIZE_K * stride_ak
+ b_ptrs += (BLOCK_SIZE_K // 2) * 16 * 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)
+ if GRID_MN <= 224:
+ tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")
+ else:
+ tl.store(c_ptrs, c, mask=c_mask)
+
+
def _e8m0_shuffle_safe(scale: torch.Tensor) -> torch.Tensor:
m, n = scale.shape
scale_padded = torch.empty(
- ((m + 255) // 256) * 256,
- ((n + 7) // 8) * 8,
- dtype=scale.dtype,
- device=scale.device,
+ ((m + 255) // 256) * 256, ((n + 7) // 8) * 8,
+ dtype=scale.dtype, device=scale.device,
)
scale_padded.fill_(0x7F)
scale_padded[:m, :n] = scale
sm, sn = scale_padded.shape
return (
scale_padded.view(sm // 32, 2, 16, sn // 8, 2, 4)
- .permute(0, 3, 5, 2, 4, 1)
- .contiguous()
- .view(sm, sn)
+ .permute(0, 3, 5, 2, 4, 1).contiguous().view(sm, sn)
)
- def _safe_wrapper(a: torch.Tensor, b_shuffle: torch.Tensor, b_scale_sh: torch.Tensor):
+ @triton.jit
+ def _reduce_m16_splitk7_kernel(
+ y_ptr, out_ptr,
+ N,
+ stride_yk, stride_ym, stride_yn,
+ stride_om, stride_on,
+ ACTUAL_KSPLIT: tl.constexpr,
+ MAX_KSPLIT: tl.constexpr,
+ BLOCK_N: tl.constexpr,
+ ):
+ pid_n = tl.program_id(axis=0)
+ offs_m = tl.arange(0, 16)
+ offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
+ n_mask = offs_n < N
+ acc = tl.zeros((16, BLOCK_N), dtype=tl.float32)
+
+ for ks in range(MAX_KSPLIT):
+ if ks < ACTUAL_KSPLIT:
+ y_ptrs = (
+ y_ptr
+ + ks * stride_yk
+ + offs_m[:, None] * stride_ym
+ + offs_n[None, :] * stride_yn
+ )
+ acc += tl.load(y_ptrs, mask=n_mask[None, :], other=0.0)
+
+ out = acc.to(out_ptr.type.element_ty)
+ out_ptrs = out_ptr + offs_m[:, None] * stride_om + offs_n[None, :] * stride_on
+ tl.store(out_ptrs, out, mask=n_mask[None, :])
+
+
+ def _safe_wrapper(a, b_shuffle, b_scale_sh):
a_q_raw, a_scale = dynamic_mxfp4_quant(a.contiguous())
a_scale_sh = _e8m0_shuffle_safe(a_scale)
return aiter.gemm_a4w4(
- a_q_raw.view(_FP4X2),
- b_shuffle,
- a_scale_sh.view(_FP8_E8M0),
- b_scale_sh,
- dtype=_BF16,
- bpreshuffle=True,
+ a_q_raw.view(_FP4X2), b_shuffle,
+ a_scale_sh.view(_FP8_E8M0), b_scale_sh,
+ dtype=_BF16, bpreshuffle=True,
)
- def _get_large_bufs(m: int, k: int, n: int, device):
- x_fp4 = torch.empty((m, k >> 1), dtype=torch.uint8, device=device)
- scale_n = (k + _QUANT_BLOCK - 1) // _QUANT_BLOCK
- scale_n_pad = ((scale_n + 7) >> 3) << 3
- scale_m_pad = ((m + 255) >> 8) << 8
- scale = torch.empty((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=device)
- padded_m = ((m + 31) >> 5) << 5
- out = torch.empty_strided(
- (padded_m, n),
- (n + _OUT_PAD_BF16, 1),
- dtype=_BF16,
- device=device,
- )
- return x_fp4, scale, scale_n, scale_n_pad, scale_m_pad, out
-
-
- def _get_route(key):
- m, n, k = key
-
- if key in _PUBLIC_SMALL:
- return ("small", _PUBLIC_SMALL[key])
- if key in _PUBLIC_LARGE:
- return ("large", _PUBLIC_LARGE[key])
-
+ def _get_route(m, n, k):
+ key = (m, n, k)
+ if key in _FUSED_CONFIGS:
+ return ("fused", _FUSED_CONFIGS[key])
+ if key in _ASM_CONFIGS:
+ return ("asm", _ASM_CONFIGS[key])
pair = (n, k)
if pair == (2112, 7168):
if m < 16:
- return ("small", _PUBLIC_TEST_SMALL[(8, 2112, 7168)])
- return ("small", _PUBLIC_SMALL[(16, 2112, 7168)])
+ return ("fused", _FUSED_K7168_KSPLIT14)
+ return ("fused", _FUSED_CONFIGS[(16, 2112, 7168)])
if pair == (3072, 1536):
- if m < 128:
- return ("small", _PUBLIC_TEST_SMALL[(16, 3072, 1536)])
- return ("large", 1)
+ if m <= 16:
+ return ("fused", _FUSED_K1536_SPLIT)
+ return ("fused", _FUSED_K1536_NOSPLIT)
if pair == (2880, 512):
if m < 16:
- return ("small", _PUBLIC_SMALL[(4, 2880, 512)])
- if m < 96:
- return ("small", _PUBLIC_SMALL[(32, 2880, 512)])
- return ("large", 2)
+ return ("fused", _FUSED_CONFIGS[(4, 2880, 512)])
+ if m < 128:
+ return ("fused", _FUSED_CONFIGS[(32, 2880, 512)])
+ return ("fused", _FUSED_BM32_K512)
if pair == (4096, 512):
- return ("small", _PUBLIC_SMALL[(32, 4096, 512)])
+ return ("fused", _FUSED_CONFIGS[(32, 4096, 512)])
if pair == (7168, 2048):
- return ("large", 2)
+ return ("fused", _FUSED_K2048_NOSPLIT)
return None
- @torch.inference_mode()
- 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 = (int(m), int(n), int(k))
- route = _get_route(key)
+ def _make_asm_handler(m, k, n, device, splitk, kernel_name=_KERNEL_32X128):
+ padded_m = ((m + 31) >> 5) << 5
+ x_fp4 = torch.empty((padded_m, k >> 1), dtype=torch.uint8, device=device)
+ sN = (k + _QUANT_BLOCK - 1) // _QUANT_BLOCK
+ sN_pad = ((sN + 7) >> 3) << 3
+ sM_pad = ((m + 255) >> 8) << 8
+ scale = torch.empty((sM_pad, sN_pad), dtype=torch.uint8, device=device)
+ out = torch.empty_strided(
+ (padded_m, n), (n + 32, 1), dtype=_BF16, device=device,
+ )
+ x_fp4_view = x_fp4.view(_FP4X2)
+ scale_view = scale.view(_FP8_E8M0)
+ q_grid = ((m + _QUANT_TILE - 1) // _QUANT_TILE, sN_pad)
+ sfp4_0 = x_fp4.stride(0)
+ sfp4_1 = x_fp4.stride(1)
+ ssc_0 = scale.stride(0)
+ ssc_1 = scale.stride(1)
+ out_slice = out[:m]
+ quant_launch = _quant_kernel_asm_layout[q_grid]
- if route is None:
- return _safe_wrapper(a, b_shuffle, b_scale_sh)
-
- route_kind, route_value = route
-
- if route_kind == "large":
- if key not in _BUFS:
- _BUFS[key] = ("large", _get_large_bufs(m, k, n, a.device))
- _, (x_fp4, scale, scale_n, scale_n_pad, scale_m_pad, out) = _BUFS[key]
- grid = ((m + _QUANT_TILE - 1) // _QUANT_TILE, scale_n_pad)
- _dynamic_mxfp4_quant_kernel_even_asm_layout[grid](
- a,
- x_fp4,
- scale,
- a.stride(0),
- a.stride(1),
- x_fp4.stride(0),
- x_fp4.stride(1),
- scale.stride(0),
- scale.stride(1),
- M=m,
- N=k,
- scaleN=scale_n,
- scaleM_pad=scale_m_pad,
- scaleN_pad=scale_n_pad,
- BLOCK_SIZE=_QUANT_TILE,
- MXFP4_QUANT_BLOCK_SIZE=_QUANT_BLOCK,
- SHUFFLE=True,
+ def handler(a, b_shuffle, b_scale_sh):
+ quant_launch(
+ a, x_fp4, scale, k, 1, sfp4_0, sfp4_1, ssc_0, ssc_1,
+ M=m, N=k, scaleN=sN,
+ scaleM_pad=sM_pad, scaleN_pad=sN_pad,
+ BLOCK_SIZE=_QUANT_TILE, MXFP4_QUANT_BLOCK_SIZE=_QUANT_BLOCK, SHUFFLE=True,
)
gemm_a4w4_asm(
- x_fp4.view(_FP4X2),
- b_shuffle,
- scale.view(_FP8_E8M0),
- b_scale_sh,
- out,
- _KERNEL_32X128,
- bpreshuffle=True,
- log2_k_split=route_value,
+ x_fp4_view, b_shuffle, scale_view, b_scale_sh,
+ out, kernel_name, bpreshuffle=True, log2_k_split=splitk,
)
- return out[:m]
+ return out_slice
- if key not in _BUFS:
- _BUFS[key] = ("small", torch.empty((m, n), dtype=_BF16, device=a.device))
- _, out = _BUFS[key]
+ return handler
- w = b_shuffle.view(torch.uint8).reshape(n // 16, k // 2 * 16)
- sm, sn = b_scale_sh.shape
- w_scales = b_scale_sh.view(torch.uint8).reshape(sm // 32, sn * 32)
- return gemm_a16wfp4_preshuffle(
- a,
- w,
- w_scales,
- prequant=True,
- y=out,
- config=route_value,
- )
+
+ def _make_fused_handler(m, n, k, device, raw_config, b_scale_sh):
+ config = dict(raw_config)
+ K_kernel = k // 2
+ if config["NUM_KSPLIT"] > 1:
+ SPLITK_BLOCK_SIZE, BLOCK_SIZE_K, NUM_KSPLIT = get_splitk(
+ K_kernel, config["BLOCK_SIZE_K"], config["NUM_KSPLIT"]
+ )
+ config["SPLITK_BLOCK_SIZE"] = SPLITK_BLOCK_SIZE
+ config["BLOCK_SIZE_K"] = BLOCK_SIZE_K
+ config["NUM_KSPLIT"] = NUM_KSPLIT
+ if config["BLOCK_SIZE_K"] >= 2 * K_kernel:
+ config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * K_kernel)
+ config["SPLITK_BLOCK_SIZE"] = 2 * K_kernel
+ config["NUM_KSPLIT"] = 1
+ config["BLOCK_SIZE_N"] = max(config["BLOCK_SIZE_N"], 32)
+ if config["NUM_KSPLIT"] <= 1:
+ config["SPLITK_BLOCK_SIZE"] = 2 * K_kernel
+
+ has_splitk = config["NUM_KSPLIT"] > 1
+ out = torch.empty((m, n), dtype=_BF16, device=device)
+ y_pp = None
+ if has_splitk:
+ y_pp = torch.empty((config["NUM_KSPLIT"], m, n), dtype=torch.float32, device=device)
+
+ BM = config["BLOCK_SIZE_M"]
+ BN = config["BLOCK_SIZE_N"]
+ total_tiles = triton.cdiv(m, BM) * triton.cdiv(n, BN)
+ grid = (config["NUM_KSPLIT"] * total_tiles,)
+ exact_large_route = "m64" if (m, n, k) == (64, 7168, 2048) else ("m256" if (m, n, k) == (256, 3072, 1536) else "")
+ stride_ck = 0 if y_pp is None else y_pp.stride(0)
+ stride_cm = out.stride(0) if y_pp is None else y_pp.stride(1)
+ stride_cn = out.stride(1) if y_pp is None else y_pp.stride(2)
+ c_ptr = y_pp if has_splitk else out
+ a_s0 = k
+ w_s0 = k // 2 * 16
+ bss_1 = b_scale_sh.size(1)
+ ws_s0 = bss_1 * 32
+
+ use_exact_m16 = has_splitk and (m, n, k) == (16, 2112, 7168)
+ if use_exact_m16:
+ fused_launch = _gemm_exact_m16_xcd_kernel[grid]
+ elif has_splitk:
+ fused_launch = _gemm_optimized_xcd_kernel[grid]
+ else:
+ if exact_large_route == "m64":
+ grid = (224,)
+ fused_launch = _enhanced_preshuffle_exact_m64_kernel[grid]
+ elif exact_large_route == "m256":
+ grid = (384,)
+ fused_launch = _enhanced_preshuffle_exact_m256_kernel[grid]
+ else:
+ fused_launch = _enhanced_preshuffle_kernel[grid]
+
+ if has_splitk:
+ ACTUAL_KSPLIT = triton.cdiv(K_kernel, config["SPLITK_BLOCK_SIZE"] // 2)
+ MAX_KSPLIT = triton.next_power_of_2(config["NUM_KSPLIT"])
+ use_specialized_reduce = (m, n, k) == (16, 2112, 7168) and MAX_KSPLIT <= 8
+ if use_specialized_reduce:
+ reduce_grid = (triton.cdiv(n, 64),)
+ else:
+ reduce_grid = (triton.cdiv(m, 16), triton.cdiv(n, 64))
+ reduce_args = (
+ m, n, y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
+ out.stride(0), out.stride(1), 16, 64, ACTUAL_KSPLIT, MAX_KSPLIT,
+ )
+ reduce_launch = _gemm_afp4wfp4_reduce_kernel[reduce_grid]
+
+ if has_splitk:
+ def handler(a, b_shuffle, b_scale_sh):
+ w = b_shuffle.view(torch.uint8)
+ ws = b_scale_sh.view(torch.uint8)
+ fused_launch(a, w, c_ptr, ws, m, n, K_kernel, a_s0, 1, w_s0, 1,
+ stride_ck, stride_cm, stride_cn, ws_s0, 1, PREQUANT=True, **config)
+ if use_specialized_reduce:
+ _reduce_m16_splitk7_kernel[reduce_grid](
+ y_pp, out, n,
+ y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
+ out.stride(0), out.stride(1),
+ ACTUAL_KSPLIT=ACTUAL_KSPLIT, MAX_KSPLIT=MAX_KSPLIT, BLOCK_N=64,
+ num_warps=4, num_stages=1,
+ )
+ else:
+ reduce_launch(y_pp, out, *reduce_args)
+ return out
+ else:
+ def handler(a, b_shuffle, b_scale_sh):
+ w = b_shuffle.view(torch.uint8)
+ ws = b_scale_sh.view(torch.uint8)
+ fused_launch(a, w, c_ptr, ws, m, n, K_kernel, a_s0, 1, w_s0, 1,
+ stride_ck, stride_cm, stride_cn, ws_s0, 1, PREQUANT=True, **config)
+ return out
+
+ return handler
+
+
+ _HANDLERS = {}
+ _last_key = None
+ _last_handler = None
+
+
+ @torch.inference_mode()
+ def custom_kernel(data: input_t) -> output_t:
+ global _last_key, _last_handler
+ a = data[0]
+ m = a.size(0)
+ k = a.size(1)
+ n = data[1].size(0)
+ key = (m, n, k)
+ if key == _last_key:
+ return _last_handler(a, data[3], data[4])
+ b_shuffle = data[3]
+ b_scale_sh = data[4]
+ if key not in _HANDLERS:
+ route = _get_route(m, n, k)
+ if route is None:
+ _HANDLERS[key] = lambda a, bs, bss: _safe_wrapper(a, bs, bss)
+ elif route[0] == "asm":
+ splitk, kname = route[1]
+ _HANDLERS[key] = _make_asm_handler(m, k, n, a.device, splitk, kname)
+ else:
+ _HANDLERS[key] = _make_fused_handler(m, n, k, a.device, route[1], b_scale_sh)
+ handler = _HANDLERS[key]
+ _last_key = key
+ _last_handler = handler
+ return handler(a, b_shuffle, b_scale_sh)
scrolls · 1131 diff lines total

Best evidence level for this revision: reported

JSON