Skip to content
KernelIndex
Search⌘K

submission 527012

_radna · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:8b3d7f7252d0d369d3b97d3a5d05bde337522edbaee1cfcd4542dfea1d0b8e5a
license declaredunknown
license concludedunknown
authors_radna
imported2026-08-26

Techniques

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

fp4Hybrid MXFP4 submission:
split-k_ASM_SPLIT_K = 0
tile-n = 64BLOCK_SIZE_N=64,

Kernel source

submission.py1455 lines
"""
Hybrid MXFP4 submission:
- Keep the proven explicit asm dispatch for five benchmark tuples and the general fallback path.
- Replace only the dominant `(16, 2112, 7168)` tuple with a vendored fix of AITER's
  broken `gemm_afp4wfp4_preshuffle` Triton kernel.
"""
from functools import lru_cache

from task import input_t, output_t

try:
    import triton
    import triton.language as tl
except Exception:
    triton = None
    tl = None

_EXPLICIT_ASM_KERNELS = {
    (4, 2880, 512): "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E",
    (16, 2112, 7168): "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
    (32, 4096, 512): "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
    (32, 2880, 512): "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
    (64, 7168, 2048): "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
    (256, 3072, 1536): "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
}
_VENDORED_TRITON_PRESHUFFLE_CONFIGS = {
    (16, 2112, 7168): {
        "BLOCK_SIZE_M": 16,
        "BLOCK_SIZE_N": 32,
        "BLOCK_SIZE_K": 1024,
        "GROUP_SIZE_M": 1,
        "NUM_KSPLIT": 7,
        "num_warps": 4,
        "num_stages": 2,
        "waves_per_eu": 4,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
    },
}
_HOT_BIGK_STATE = {"enabled": True, "error": None}
_A16WFP4_PRESHUFFLE_STATE = {"enabled": True, "error": None}
_FUSED_A16WFP4_SHAPES = {
    (4, 2880, 512),
    (32, 4096, 512),
    (32, 2880, 512),
}
_ASM_SPLIT_K = 0
_EXPLICIT_ASM_LOG2_K_SPLIT = {}
_WORKSPACE_CACHE = {}
_WORKSPACE_CACHE_LIMIT = 16


def _cdiv(x, y):
    return (x + y - 1) // y


def _next_power_of_two(x):
    return 1 if x <= 1 else 1 << (x - 1).bit_length()


def _get_workspace_tensor(tag, reference, shape, dtype):
    import torch

    device = reference.device
    device_index = -1 if device.index is None else device.index
    key = (tag, device.type, device_index, str(dtype), tuple(shape))
    cached = _WORKSPACE_CACHE.get(key)
    if cached is not None:
        return cached
    if len(_WORKSPACE_CACHE) >= _WORKSPACE_CACHE_LIMIT:
        _WORKSPACE_CACHE.clear()
    cached = torch.empty(shape, dtype=dtype, device=device)
    _WORKSPACE_CACHE[key] = cached
    return cached


def _get_splitk(K, block_size_k, num_ksplit):
    num_ksplit_step = 2
    block_size_k_step = 2
    splitk_block_size = _cdiv(2 * _cdiv(K, num_ksplit), block_size_k) * block_size_k
    while num_ksplit > 1 and block_size_k > 16:
        if (
            K % (splitk_block_size // 2) == 0
            and splitk_block_size % block_size_k == 0
            and K % (block_size_k // 2) == 0
        ):
            break
        if K % (splitk_block_size // 2) != 0 and num_ksplit > 1:
            num_ksplit //= num_ksplit_step
        elif splitk_block_size % block_size_k != 0:
            if num_ksplit > 1:
                num_ksplit //= num_ksplit_step
            elif block_size_k > 16:
                block_size_k //= block_size_k_step
        elif K % (block_size_k // 2) != 0 and block_size_k > 16:
            block_size_k //= block_size_k_step
        else:
            break
        splitk_block_size = _cdiv(2 * _cdiv(K, num_ksplit), block_size_k) * block_size_k
    num_ksplit = _cdiv(K, splitk_block_size // 2)
    return splitk_block_size, block_size_k, num_ksplit


if triton is not None:

    @triton.jit
    def _remap_xcd(pid, GRID_MN, NUM_XCDS: tl.constexpr = 8):
        pids_per_xcd = (GRID_MN + NUM_XCDS - 1) // NUM_XCDS
        tall_xcds = GRID_MN % NUM_XCDS
        tall_xcds = NUM_XCDS if tall_xcds == 0 else tall_xcds
        xcd = pid % NUM_XCDS
        local_pid = pid // NUM_XCDS
        if xcd < tall_xcds:
            pid = xcd * pids_per_xcd + local_pid
        else:
            pid = (
                tall_xcds * pids_per_xcd
                + (xcd - tall_xcds) * (pids_per_xcd - 1)
                + local_pid
            )
        return pid


    @triton.jit
    def _pid_grid(pid: int, num_pid_m: int, num_pid_n: int, GROUP_SIZE_M: tl.constexpr = 1):
        if GROUP_SIZE_M == 1:
            pid_m = pid // num_pid_n
            pid_n = pid % num_pid_n
        else:
            num_pid_in_group = GROUP_SIZE_M * num_pid_n
            group_id = pid // num_pid_in_group
            first_pid_m = group_id * GROUP_SIZE_M
            group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M)
            tl.assume(group_size_m >= 0)
            pid_m = first_pid_m + (pid % group_size_m)
            pid_n = (pid % num_pid_in_group) // group_size_m
        return pid_m, pid_n


    @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),
        }
    )
    @triton.jit
    def _vendored_gemm_afp4wfp4_preshuffle_kernel(
        a_ptr,
        b_ptr,
        c_ptr,
        a_scales_ptr,
        b_scales_ptr,
        M,
        N,
        K,
        stride_am,
        stride_ak,
        stride_bn,
        stride_bk,
        stride_ck,
        stride_cm,
        stride_cn,
        stride_asm,
        stride_ask,
        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,
        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_asm > 0)
        tl.assume(stride_ask > 0)
        tl.assume(stride_bsk > 0)
        tl.assume(stride_bsn > 0)

        grid_mn = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)
        pid_unified = tl.program_id(axis=0)
        pid_unified = _remap_xcd(pid_unified, grid_mn * NUM_KSPLIT, NUM_XCDS=8)
        pid_k = pid_unified % NUM_KSPLIT
        pid = pid_unified // NUM_KSPLIT
        num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
        num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)

        if NUM_KSPLIT == 1:
            pid_m, pid_n = _pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
        else:
            pid_m = pid // num_pid_n
            pid_n = pid % num_pid_n

        tl.assume(pid_m >= 0)
        tl.assume(pid_n >= 0)
        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 = tl.arange(0, BLOCK_SIZE_K // 2)
            offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
            offs_k_split = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_k
            offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr

            offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
            offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
            a_ptrs = a_ptr + (
                offs_am[:, None] * stride_am + offs_k_split[None, :] * stride_ak
            )
            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
            )

            if BLOCK_SIZE_M < 32:
                offs_ks_non_shuffle = (
                    pid_k * (SPLITK_BLOCK_SIZE // scale_group_size)
                ) + tl.arange(0, BLOCK_SIZE_K // scale_group_size)
                a_scale_ptrs = (
                    a_scales_ptr
                    + offs_am[:, None] * stride_asm
                    + offs_ks_non_shuffle[None, :] * stride_ask
                )
            else:
                offs_asm = (
                    pid_m * (BLOCK_SIZE_M // 32) + tl.arange(0, (BLOCK_SIZE_M // 32))
                ) % M
                a_scale_ptrs = (
                    a_scales_ptr
                    + offs_asm[:, None] * stride_asm
                    + offs_ks[None, :] * stride_ask
                )

            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):
                if BLOCK_SIZE_M < 32:
                    a_scales = tl.load(a_scale_ptrs)
                else:
                    a_scales = (
                        tl.load(a_scale_ptrs)
                        .reshape(
                            BLOCK_SIZE_M // 32,
                            BLOCK_SIZE_K // scale_group_size // 8,
                            4,
                            16,
                            2,
                            2,
                            1,
                        )
                        .permute(0, 5, 3, 1, 4, 2, 6)
                        .reshape(BLOCK_SIZE_M, BLOCK_SIZE_K // scale_group_size)
                    )

                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 = tl.load(a_ptrs)
                    b = tl.load(b_ptrs, cache_modifier=cache_modifier)
                else:
                    a = tl.load(
                        a_ptrs,
                        mask=offs_k[None, :] < K - k * (BLOCK_SIZE_K // 2),
                        other=0,
                    )
                    b = tl.load(
                        b_ptrs,
                        mask=offs_k_shuffle_arr[None, :] < (K - k * (BLOCK_SIZE_K // 2)) * 16,
                        other=0,
                        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)
                )

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

                a_ptrs += (BLOCK_SIZE_K // 2) * stride_ak
                b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
                if BLOCK_SIZE_M < 32:
                    a_scale_ptrs += (BLOCK_SIZE_K // scale_group_size) * stride_ask
                else:
                    a_scale_ptrs += BLOCK_SIZE_K * stride_ask
                b_scale_ptrs += BLOCK_SIZE_K * stride_bsk

            c = accumulator.to(c_ptr.type.element_ty)
            offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
            offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
            c_ptrs = (
                c_ptr
                + stride_cm * offs_cm[:, None]
                + stride_cn * offs_cn[None, :]
                + pid_k * stride_ck
            )
            c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
            tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")


    @triton.jit
    def _vendored_gemm_afp4wfp4_reduce_kernel(
        c_in_ptr,
        c_out_ptr,
        M,
        N,
        stride_c_in_k,
        stride_c_in_m,
        stride_c_in_n,
        stride_c_out_m,
        stride_c_out_n,
        BLOCK_SIZE_M: tl.constexpr,
        BLOCK_SIZE_N: tl.constexpr,
        ACTUAL_KSPLIT: tl.constexpr,
        MAX_KSPLIT: tl.constexpr,
    ):
        pid_m = tl.program_id(axis=0)
        pid_n = tl.program_id(axis=1)

        offs_m = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
        offs_n = (pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)) % N
        offs_k = tl.arange(0, MAX_KSPLIT)
        c_in_ptrs = (
            c_in_ptr
            + (offs_k[:, None, None] * stride_c_in_k)
            + (offs_m[None, :, None] * stride_c_in_m)
            + (offs_n[None, None, :] * stride_c_in_n)
        )

        if ACTUAL_KSPLIT == MAX_KSPLIT:
            c = tl.load(c_in_ptrs)
        else:
            c = tl.load(c_in_ptrs, mask=offs_k[:, None, None] < ACTUAL_KSPLIT)
        c = tl.sum(c, axis=0)
        c = c.to(c_out_ptr.type.element_ty)

        c_out_ptrs = (
            c_out_ptr
            + (offs_m[:, None] * stride_c_out_m)
            + (offs_n[None, :] * stride_c_out_n)
        )
        tl.store(c_out_ptrs, c)


    @triton.jit
    def _vendored_gemm_afp4wfp4_preshuffle_hot_exact_kernel(
        a_ptr,
        b_ptr,
        c_ptr,
        a_scales_ptr,
        b_scales_ptr,
        stride_am,
        stride_ak,
        stride_bn,
        stride_bk,
        stride_ck,
        stride_cm,
        stride_cn,
        stride_asm,
        stride_ask,
        stride_bsn,
        stride_bsk,
        num_warps: tl.constexpr,
        num_stages: tl.constexpr,
        waves_per_eu: tl.constexpr,
        matrix_instr_nonkdim: tl.constexpr,
        cache_modifier: tl.constexpr,
    ):
        pid_unified = tl.program_id(axis=0)
        pid_unified = _remap_xcd(pid_unified, 66 * 7, NUM_XCDS=8)
        pid_k = pid_unified % 7
        pid_n = pid_unified // 7

        offs_am = tl.arange(0, 16)
        offs_bn = pid_n * 2 + tl.arange(0, 2)
        offs_k = tl.arange(0, 256)
        offs_k_shuffle_arr = tl.arange(0, 4096)

        a_ptrs0 = a_ptr + offs_am[:, None] * stride_am + (pid_k * 512 + offs_k)[None, :] * stride_ak
        a_ptrs1 = a_ptr + offs_am[:, None] * stride_am + (pid_k * 512 + 256 + offs_k)[None, :] * stride_ak
        b_ptrs0 = b_ptr + offs_bn[:, None] * stride_bn + (pid_k * 8192 + offs_k_shuffle_arr)[None, :] * stride_bk
        b_ptrs1 = b_ptr + offs_bn[:, None] * stride_bn + (pid_k * 8192 + 4096 + offs_k_shuffle_arr)[None, :] * stride_bk

        offs_bsn = pid_n + tl.arange(0, 1)
        b_scale_ptrs0 = (
            b_scales_ptr
            + offs_bsn[:, None] * stride_bsn
            + (pid_k * 1024 + tl.arange(0, 512))[None, :] * stride_bsk
        )
        b_scale_ptrs1 = (
            b_scales_ptr
            + offs_bsn[:, None] * stride_bsn
            + (pid_k * 1024 + 512 + tl.arange(0, 512))[None, :] * stride_bsk
        )

        a_scale_ptrs0 = (
            a_scales_ptr
            + offs_am[:, None] * stride_asm
            + (pid_k * 32 + tl.arange(0, 16))[None, :] * stride_ask
        )
        a_scale_ptrs1 = (
            a_scales_ptr
            + offs_am[:, None] * stride_asm
            + (pid_k * 32 + 16 + tl.arange(0, 16))[None, :] * stride_ask
        )

        accumulator = tl.zeros((16, 32), dtype=tl.float32)

        a_scales0 = tl.load(a_scale_ptrs0)
        b_scales0 = (
            tl.load(b_scale_ptrs0, cache_modifier=cache_modifier)
            .reshape(1, 2, 4, 16, 2, 2, 1)
            .permute(0, 5, 3, 1, 4, 2, 6)
            .reshape(32, 16)
        )
        a0 = tl.load(a_ptrs0)
        b0 = tl.load(b_ptrs0, cache_modifier=cache_modifier)
        b0 = (
            b0.reshape(1, 2, 8, 2, 16, 16)
            .permute(0, 1, 4, 2, 3, 5)
            .reshape(32, 256)
            .trans(1, 0)
        )
        accumulator = tl.dot_scaled(
            a0, a_scales0, "e2m1", b0, b_scales0, "e2m1", accumulator
        )

        a_scales1 = tl.load(a_scale_ptrs1)
        b_scales1 = (
            tl.load(b_scale_ptrs1, cache_modifier=cache_modifier)
            .reshape(1, 2, 4, 16, 2, 2, 1)
            .permute(0, 5, 3, 1, 4, 2, 6)
            .reshape(32, 16)
        )
        a1 = tl.load(a_ptrs1)
        b1 = tl.load(b_ptrs1, cache_modifier=cache_modifier)
        b1 = (
            b1.reshape(1, 2, 8, 2, 16, 16)
            .permute(0, 1, 4, 2, 3, 5)
            .reshape(32, 256)
            .trans(1, 0)
        )
        accumulator = tl.dot_scaled(
            a1, a_scales1, "e2m1", b1, b_scales1, "e2m1", accumulator
        )

        offs_cn = pid_n * 32 + tl.arange(0, 32).to(tl.int64)
        c_ptrs = (
            c_ptr
            + pid_k * stride_ck
            + stride_cm * offs_am[:, None]
            + stride_cn * offs_cn[None, :]
        )
        tl.store(c_ptrs, accumulator.to(c_ptr.type.element_ty))

    @triton.jit
    def _vendored_gemm_afp4wfp4_preshuffle_hot_exact_k1024_kernel(
        a_ptr,
        b_ptr,
        c_ptr,
        a_scales_ptr,
        b_scales_ptr,
        stride_am,
        stride_ak,
        stride_bn,
        stride_bk,
        stride_ck,
        stride_cm,
        stride_cn,
        stride_asm,
        stride_ask,
        stride_bsn,
        stride_bsk,
        num_warps: tl.constexpr,
        num_stages: tl.constexpr,
        waves_per_eu: tl.constexpr,
        matrix_instr_nonkdim: tl.constexpr,
        cache_modifier: tl.constexpr,
    ):
        # Derive the 1024-element packing directly from the known-correct generic preshuffle kernel:
        # - K is byte-addressed for fp4x2 tensors (two fp4 values per byte)
        # - For the hot tuple, each split covers 1024 fp4 elements = 512 bytes.
        pid_unified = tl.program_id(axis=0)
        pid_unified = _remap_xcd(pid_unified, 66 * 7, NUM_XCDS=8)
        pid_k = pid_unified % 7
        pid_n = pid_unified // 7

        offs_am = tl.arange(0, 16)
        offs_bn = pid_n * 2 + tl.arange(0, 2)
        offs_k = tl.arange(0, 512)
        offs_k_shuffle_arr = tl.arange(0, 8192)

        a_ptrs = (
            a_ptr
            + offs_am[:, None] * stride_am
            + (pid_k * 512 + offs_k)[None, :] * stride_ak
        )
        b_ptrs = (
            b_ptr
            + offs_bn[:, None] * stride_bn
            + (pid_k * 8192 + offs_k_shuffle_arr)[None, :] * stride_bk
        )

        offs_bsn = pid_n + tl.arange(0, 1)
        b_scale_ptrs = (
            b_scales_ptr
            + offs_bsn[:, None] * stride_bsn
            + (pid_k * 1024 + tl.arange(0, 1024))[None, :] * stride_bsk
        )
        a_scale_ptrs = (
            a_scales_ptr
            + offs_am[:, None] * stride_asm
            + (pid_k * 32 + tl.arange(0, 32))[None, :] * stride_ask
        )

        accumulator = tl.zeros((16, 32), dtype=tl.float32)

        a_scales = tl.load(a_scale_ptrs)
        b_scales = (
            tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
            .reshape(1, 4, 4, 16, 2, 2, 1)
            .permute(0, 5, 3, 1, 4, 2, 6)
            .reshape(32, 32)
        )
        a = tl.load(a_ptrs)
        b = tl.load(b_ptrs, cache_modifier=cache_modifier)
        b = (
            b.reshape(1, 2, 16, 2, 16, 16)
            .permute(0, 1, 4, 2, 3, 5)
            .reshape(32, 512)
            .trans(1, 0)
        )
        accumulator = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", accumulator)

        offs_cn = pid_n * 32 + tl.arange(0, 32).to(tl.int64)
        c_ptrs = (
            c_ptr
            + pid_k * stride_ck
            + stride_cm * offs_am[:, None]
            + stride_cn * offs_cn[None, :]
        )
        tl.store(c_ptrs, accumulator.to(c_ptr.type.element_ty))


    @triton.jit
    def _vendored_gemm_afp4wfp4_reduce_hot_exact_kernel(
        c_in_ptr,
        c_out_ptr,
        stride_c_in_k,
        stride_c_in_m,
        stride_c_in_n,
        stride_c_out_m,
        stride_c_out_n,
        BLOCK_SIZE_N: tl.constexpr,
    ):
        pid_n = tl.program_id(axis=0)
        offs_m = tl.arange(0, 16)
        offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
        c0 = tl.load(
            c_in_ptr
            + 0 * stride_c_in_k
            + offs_m[:, None] * stride_c_in_m
            + offs_n[None, :] * stride_c_in_n
        )
        c1 = tl.load(
            c_in_ptr
            + 1 * stride_c_in_k
            + offs_m[:, None] * stride_c_in_m
            + offs_n[None, :] * stride_c_in_n
        )
        c2 = tl.load(
            c_in_ptr
            + 2 * stride_c_in_k
            + offs_m[:, None] * stride_c_in_m
            + offs_n[None, :] * stride_c_in_n
        )
        c3 = tl.load(
            c_in_ptr
            + 3 * stride_c_in_k
            + offs_m[:, None] * stride_c_in_m
            + offs_n[None, :] * stride_c_in_n
        )
        c4 = tl.load(
            c_in_ptr
            + 4 * stride_c_in_k
            + offs_m[:, None] * stride_c_in_m
            + offs_n[None, :] * stride_c_in_n
        )
        c5 = tl.load(
            c_in_ptr
            + 5 * stride_c_in_k
            + offs_m[:, None] * stride_c_in_m
            + offs_n[None, :] * stride_c_in_n
        )
        c6 = tl.load(
            c_in_ptr
            + 6 * stride_c_in_k
            + offs_m[:, None] * stride_c_in_m
            + offs_n[None, :] * stride_c_in_n
        )
        c = (c0 + c1 + c2 + c3 + c4 + c5 + c6).to(c_out_ptr.type.element_ty)
        c_out_ptrs = (
            c_out_ptr
            + offs_m[:, None] * stride_c_out_m
            + offs_n[None, :] * stride_c_out_n
        )
        tl.store(c_out_ptrs, c)

    @triton.jit
    def _mxfp4_quant_contract_op(
        x,
        BLOCK_SIZE_N: tl.constexpr,
        BLOCK_SIZE_M: tl.constexpr,
        MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
    ):
        # Contract-aligned MXFP4 quantization op (matches fp4_utils conversion).
        #
        # Inputs:
        #   x: [BLOCK_SIZE_M, BLOCK_SIZE_N] bf16/fp16/fp32
        # Returns:
        #   x_fp4: [BLOCK_SIZE_M, BLOCK_SIZE_N // 2] uint8 (packed e2m1 fp4x2)
        #   bs_e8m0: [BLOCK_SIZE_M, BLOCK_SIZE_N // 32] uint8 (e8m0, biased by 127)
        x = x.to(tl.float32)
        num_quant_blocks: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
        x = x.reshape(BLOCK_SIZE_M, num_quant_blocks, MXFP4_QUANT_BLOCK_SIZE)

        amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
        amax = amax.to(tl.int32, bitcast=True)
        amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
        amax = amax.to(tl.float32, bitcast=True)
        scale_e8m0_unbiased = tl.log2(amax).floor() - 2
        scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
        quant_scale = tl.exp2(-scale_e8m0_unbiased)

        qx = x * quant_scale
        bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127

        qx = qx.to(tl.uint32, bitcast=True)
        s = qx & 0x80000000
        e = (qx >> 23) & 0xFF
        m = qx & 0x7FFFFF

        e8_bias: tl.constexpr = 127
        e2_bias: tl.constexpr = 1

        adjusted_exponents = tl.core.sub(e8_bias, e + 1, sanitize_overflow=False)
        m = tl.where(e < e8_bias, (0x400000 | (m >> 1)) >> adjusted_exponents, m)
        e = tl.maximum(e, e8_bias - e2_bias) - (e8_bias - e2_bias)

        e2m1_tmp = tl.minimum((((e << 2) | (m >> 21)) + 1) >> 1, 0x7)
        e2m1_value = ((s >> 28) | e2m1_tmp).to(tl.uint8)

        e2m1_value = tl.reshape(
            e2m1_value,
            [BLOCK_SIZE_M, num_quant_blocks, MXFP4_QUANT_BLOCK_SIZE // 2, 2],
        )
        evens, odds = tl.split(e2m1_value)
        x_fp4 = evens | (odds << 4)
        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.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),
        }
    )
    @triton.jit
    def _vendored_gemm_a16wfp4_preshuffle_contract_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,
        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)

        grid_mn = tl.cdiv(M, BLOCK_SIZE_M) * tl.cdiv(N, BLOCK_SIZE_N)
        pid_unified = tl.program_id(axis=0)
        pid_unified = _remap_xcd(pid_unified, grid_mn * NUM_KSPLIT, NUM_XCDS=8)
        pid_k = pid_unified % NUM_KSPLIT
        pid = pid_unified // NUM_KSPLIT
        num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
        num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)

        if NUM_KSPLIT == 1:
            pid_m, pid_n = _pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
        else:
            pid_m = pid // num_pid_n
            pid_n = pid % num_pid_n

        tl.assume(pid_m >= 0)
        tl.assume(pid_n >= 0)

        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(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)
                    b = tl.load(b_ptrs, cache_modifier=cache_modifier)
                else:
                    a_bf16 = tl.load(
                        a_ptrs,
                        mask=offs_k_bf16[None, :] < 2 * K,
                        other=0,
                    )
                    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_contract_op(
                    a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, scale_group_size
                )
                accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")

                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)
            tl.store(c_ptrs, c, mask=c_mask)


@lru_cache(maxsize=1)
def _get_quant_func():
    import aiter
    from aiter import QuantType

    return aiter.get_triton_quant(QuantType.per_1x32)


@lru_cache(maxsize=1)
def _get_fp4_utils():
    from aiter.utility import fp4_utils

    return fp4_utils


@lru_cache(maxsize=128)
def _get_a16wfp4_preshuffle_config(M, N, K_bytes):
    # Reuse AITER's shape-based config selection (tuned when available).
    from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import _get_config

    config, _tuned = _get_config(M, N, K_bytes, True)
    return config


def _quant_block_size_m(m):
    if m <= 8:
        return 8
    if m <= 16:
        return 16
    if m <= 32:
        return 32
    if m <= 64:
        return 64
    return 128


def _shape_tuned_quant(A, shuffle):
    import torch
    from aiter import dtypes

    if triton is None:
        return _get_quant_func()(A, shuffle=shuffle)

    fp4_utils = _get_fp4_utils()
    M, N = A.shape

    # Preserve AITER's output contract while shrinking row tiles for small-M cases.
    x_fp4 = torch.empty((M, N // 2), dtype=torch.uint8, device=A.device)
    scale_n_valid = _cdiv(N, 32)
    scale_n_pad = _cdiv(scale_n_valid, 8) * 8
    scale_m_pad = _cdiv(M, 32) * 32
    blockscale_e8m0 = torch.empty(
        (_cdiv(M, 256) * 256, scale_n_pad),
        dtype=torch.uint8,
        device=A.device,
    )
    block_size_m = _quant_block_size_m(M)
    grid = (_cdiv(M, block_size_m), scale_n_pad)

    fp4_utils._dynamic_mxfp4_quant_kernel_asm_layout[grid](
        A,
        x_fp4,
        blockscale_e8m0,
        *A.stride(),
        *x_fp4.stride(),
        *blockscale_e8m0.stride(),
        M=M,
        N=N,
        scaleN=scale_n_valid,
        scaleM_pad=scale_m_pad,
        scaleN_pad=scale_n_pad,
        BLOCK_SIZE=block_size_m,
        MXFP4_QUANT_BLOCK_SIZE=32,
        SCALING_MODE=0,
        SHUFFLE=shuffle,
    )
    return (x_fp4.view(dtypes.fp4x2), blockscale_e8m0.view(dtypes.fp8_e8m0))


def _shape_tuned_quant_hot_exact(A):
    import torch
    from aiter import dtypes

    if triton is None:
        return _shape_tuned_quant(A, shuffle=False)

    fp4_utils = _get_fp4_utils()
    M, N = A.shape
    if (M, N) != (16, 7168):
        raise RuntimeError(f"unexpected hot exact quant shape {(M, N)}")

    x_fp4 = _get_workspace_tensor("hot_exact_x_fp4", A, (16, 7168 // 2), torch.uint8)
    blockscale_e8m0 = _get_workspace_tensor(
        "hot_exact_scale_e8m0", A, (16, 7168 // 32), torch.uint8
    )

    grid = (1, 7168 // 32)
    fp4_utils._dynamic_mxfp4_quant_kernel_asm_layout[grid](
        A,
        x_fp4,
        blockscale_e8m0,
        *A.stride(),
        *x_fp4.stride(),
        *blockscale_e8m0.stride(),
        M=16,
        N=7168,
        scaleN=7168 // 32,
        scaleM_pad=32,
        scaleN_pad=7168 // 32,
        BLOCK_SIZE=16,
        MXFP4_QUANT_BLOCK_SIZE=32,
        SCALING_MODE=0,
        SHUFFLE=False,
    )
    return (x_fp4.view(dtypes.fp4x2), blockscale_e8m0.view(dtypes.fp8_e8m0))


def _run_explicit_asm(
    aiter,
    A_q,
    B_shuffle,
    A_scale_sh,
    B_scale_sh,
    m,
    n,
    kernel_name,
    log2_k_split,
):
    import torch
    from aiter import dtypes

    out = torch.empty(((m + 31) // 32 * 32, n), dtype=dtypes.bf16, device=A_q.device)
    aiter.gemm_a4w4_asm(
        A_q.view(m, A_q.shape[-1]),
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        out,
        kernel_name,
        None,
        1.0,
        0.0,
        True,
        log2_k_split=log2_k_split,
    )
    return out[:m]


def _reshape_preshuffle_weight(weight, n, k):
    import torch

    weight_u8 = weight.view(torch.uint8)
    assert weight_u8.shape == (n, k // 2), (
        f"expected preshuffled weight shape {(n, k // 2)}, got {tuple(weight_u8.shape)}"
    )
    return weight_u8.view(n // 16, weight_u8.shape[1] * 16)


def _reshape_preshuffle_scales(scales, k):
    import torch

    scales_u8 = scales.view(torch.uint8)
    assert scales_u8.shape[0] % 32 == 0, (
        f"expected padded scale rows divisible by 32, got {scales_u8.shape[0]}"
    )
    assert scales_u8.numel() % k == 0, (
        f"expected scale bytes divisible by K={k}, got {scales_u8.numel()}"
    )
    return scales_u8.view(scales_u8.shape[0] // 32, scales_u8.shape[1] * 32)


def _run_vendored_triton_a16wfp4_preshuffle_contract(A, B_shuffle, B_scale_sh):
    import torch
    from aiter import dtypes

    if triton is None:
        raise RuntimeError("triton is required for the fused BF16xFP4 preshuffle path")

    M, K_full = A.shape
    N = B_shuffle.shape[0]
    if K_full % 32 != 0:
        raise RuntimeError(f"unsupported K={K_full} for MXFP4 group=32")
    if K_full % 2 != 0:
        raise RuntimeError(f"unsupported odd K={K_full} for fp4x2 packing")

    K_bytes = K_full // 2
    packed_w = _reshape_preshuffle_weight(B_shuffle, N, K_full)
    packed_scales = _reshape_preshuffle_scales(B_scale_sh, K_full)

    local_config = dict(_get_a16wfp4_preshuffle_config(M, N, K_bytes))
    return_y_pp = local_config["NUM_KSPLIT"] > 1

    splitk_block_size, block_size_k, num_ksplit = _get_splitk(
        K_bytes,
        local_config["BLOCK_SIZE_K"],
        local_config["NUM_KSPLIT"],
    )
    local_config["SPLITK_BLOCK_SIZE"] = splitk_block_size
    local_config["BLOCK_SIZE_K"] = block_size_k
    local_config["NUM_KSPLIT"] = num_ksplit

    if local_config["BLOCK_SIZE_K"] >= 2 * K_bytes:
        local_config["BLOCK_SIZE_K"] = _next_power_of_two(2 * K_bytes)
        local_config["SPLITK_BLOCK_SIZE"] = 2 * K_bytes
        local_config["NUM_KSPLIT"] = 1
        return_y_pp = False

    local_config["BLOCK_SIZE_N"] = max(local_config["BLOCK_SIZE_N"], 32)

    y = torch.empty((M, N), dtype=dtypes.bf16, device=A.device)
    if return_y_pp:
        y_pp = torch.empty(
            (local_config["NUM_KSPLIT"], M, N),
            dtype=torch.float32,
            device=A.device,
        )
    else:
        y_pp = None

    grid = lambda META: (  # noqa: E731
        (
            META["NUM_KSPLIT"]
            * triton.cdiv(M, META["BLOCK_SIZE_M"])
            * triton.cdiv(N, META["BLOCK_SIZE_N"])
        ),
    )

    _vendored_gemm_a16wfp4_preshuffle_contract_kernel[grid](
        A,
        packed_w,
        y if y_pp is None else y_pp,
        packed_scales,
        M,
        N,
        K_bytes,
        A.stride(0),
        A.stride(1),
        packed_w.stride(0),
        packed_w.stride(1),
        0 if y_pp is None else y_pp.stride(0),
        y.stride(0) if y_pp is None else y_pp.stride(1),
        y.stride(1) if y_pp is None else y_pp.stride(2),
        packed_scales.stride(0),
        packed_scales.stride(1),
        **local_config,
    )

    if y_pp is None:
        return y

    reduce_block_size_m = 16
    reduce_block_size_n = 64
    actual_ksplit = triton.cdiv(K_bytes, local_config["SPLITK_BLOCK_SIZE"] // 2)
    grid_reduce = (
        triton.cdiv(M, reduce_block_size_m),
        triton.cdiv(N, reduce_block_size_n),
    )
    _vendored_gemm_afp4wfp4_reduce_kernel[grid_reduce](
        y_pp,
        y,
        M,
        N,
        y_pp.stride(0),
        y_pp.stride(1),
        y_pp.stride(2),
        y.stride(0),
        y.stride(1),
        reduce_block_size_m,
        reduce_block_size_n,
        actual_ksplit,
        _next_power_of_two(local_config["NUM_KSPLIT"]),
    )
    return y


def _run_vendored_triton_preshuffle(A_q, B_shuffle, A_scale, B_scale_sh, config):
    import torch
    from aiter import dtypes

    if triton is None:
        raise RuntimeError("triton is required for the vendored preshuffle path")

    M, K = A_q.shape
    N = B_shuffle.shape[0]
    K = K

    packed_w = _reshape_preshuffle_weight(B_shuffle, N, K * 2)
    packed_scales = _reshape_preshuffle_scales(B_scale_sh, K * 2)

    local_config = dict(config)
    return_y_pp = local_config["NUM_KSPLIT"] > 1

    splitk_block_size, block_size_k, num_ksplit = _get_splitk(
        K,
        local_config["BLOCK_SIZE_K"],
        local_config["NUM_KSPLIT"],
    )
    local_config["SPLITK_BLOCK_SIZE"] = splitk_block_size
    local_config["BLOCK_SIZE_K"] = block_size_k
    local_config["NUM_KSPLIT"] = num_ksplit

    if local_config["BLOCK_SIZE_K"] >= 2 * K:
        local_config["BLOCK_SIZE_K"] = _next_power_of_two(2 * K)
        local_config["SPLITK_BLOCK_SIZE"] = 2 * K
        local_config["NUM_KSPLIT"] = 1
        return_y_pp = False

    local_config["BLOCK_SIZE_N"] = max(local_config["BLOCK_SIZE_N"], 32)

    y = torch.empty((M, N), dtype=dtypes.bf16, device=A_q.device)
    if return_y_pp:
        y_pp = torch.empty(
            (local_config["NUM_KSPLIT"], M, N),
            dtype=torch.float32,
            device=A_q.device,
        )
    else:
        y_pp = None

    grid = lambda META: (  # noqa: E731
        (
            META["NUM_KSPLIT"]
            * triton.cdiv(M, META["BLOCK_SIZE_M"])
            * triton.cdiv(N, META["BLOCK_SIZE_N"])
        ),
    )

    _vendored_gemm_afp4wfp4_preshuffle_kernel[grid](
        A_q.view(torch.uint8),
        packed_w,
        y if y_pp is None else y_pp,
        A_scale.view(torch.uint8),
        packed_scales,
        M,
        N,
        K,
        A_q.stride(0),
        A_q.stride(1),
        packed_w.stride(0),
        packed_w.stride(1),
        0 if y_pp is None else y_pp.stride(0),
        y.stride(0) if y_pp is None else y_pp.stride(1),
        y.stride(1) if y_pp is None else y_pp.stride(2),
        A_scale.stride(0),
        A_scale.stride(1),
        packed_scales.stride(0),
        packed_scales.stride(1),
        **local_config,
    )

    if y_pp is None:
        return y

    reduce_block_size_m = 16
    reduce_block_size_n = 64
    actual_ksplit = triton.cdiv(K, local_config["SPLITK_BLOCK_SIZE"] // 2)
    grid_reduce = (
        triton.cdiv(M, reduce_block_size_m),
        triton.cdiv(N, reduce_block_size_n),
    )
    _vendored_gemm_afp4wfp4_reduce_kernel[grid_reduce](
        y_pp,
        y,
        M,
        N,
        y_pp.stride(0),
        y_pp.stride(1),
        y_pp.stride(2),
        y.stride(0),
        y.stride(1),
        reduce_block_size_m,
        reduce_block_size_n,
        actual_ksplit,
        _next_power_of_two(local_config["NUM_KSPLIT"]),
    )
    return y


def _run_vendored_triton_preshuffle_hot_exact(A_q, B_shuffle, A_scale, B_scale_sh, config):
    import torch
    from aiter import dtypes

    if triton is None:
        raise RuntimeError("triton is required for the vendored preshuffle hot path")

    M, K = A_q.shape
    N = B_shuffle.shape[0]
    if (M, N, K) != (16, 2112, 3584):
        raise RuntimeError(f"unexpected hot exact shape {(M, N, K)}")

    packed_w = _reshape_preshuffle_weight(B_shuffle, N, K * 2)
    packed_scales = _reshape_preshuffle_scales(B_scale_sh, K * 2)

    local_config = dict(config)
    splitk_block_size, block_size_k, num_ksplit = _get_splitk(
        K,
        local_config["BLOCK_SIZE_K"],
        local_config["NUM_KSPLIT"],
    )
    local_config["SPLITK_BLOCK_SIZE"] = splitk_block_size
    local_config["BLOCK_SIZE_K"] = block_size_k
    local_config["NUM_KSPLIT"] = num_ksplit

    if not (
        local_config["BLOCK_SIZE_M"] == 16
        and local_config["BLOCK_SIZE_N"] == 32
        and local_config["BLOCK_SIZE_K"] in (512, 1024)
        and local_config["NUM_KSPLIT"] == 7
    ):
        return _run_vendored_triton_preshuffle(A_q, B_shuffle, A_scale, B_scale_sh, config)

    y_pp = _get_workspace_tensor("hot_exact_y_pp", A_q, (7, 16, 2112), torch.float32)
    y = torch.empty((16, 2112), dtype=dtypes.bf16, device=A_q.device)

    grid = (66 * 7,)
    if local_config["BLOCK_SIZE_K"] == 1024 and _HOT_BIGK_STATE["enabled"]:
        try:
            _vendored_gemm_afp4wfp4_preshuffle_hot_exact_k1024_kernel[grid](
                A_q.view(torch.uint8),
                packed_w,
                y_pp,
                A_scale.view(torch.uint8),
                packed_scales,
                A_q.stride(0),
                A_q.stride(1),
                packed_w.stride(0),
                packed_w.stride(1),
                y_pp.stride(0),
                y_pp.stride(1),
                y_pp.stride(2),
                A_scale.stride(0),
                A_scale.stride(1),
                packed_scales.stride(0),
                packed_scales.stride(1),
                num_warps=local_config["num_warps"],
                num_stages=local_config["num_stages"],
                waves_per_eu=local_config["waves_per_eu"],
                matrix_instr_nonkdim=local_config["matrix_instr_nonkdim"],
                cache_modifier=local_config["cache_modifier"],
            )
        except Exception as exc:
            # One-shot disable: if the 1024-kernel compilation fails in the runner, keep the
            # canonical 512-kernel behavior for the remainder of the process.
            _HOT_BIGK_STATE["enabled"] = False
            _HOT_BIGK_STATE["error"] = repr(exc)

            _vendored_gemm_afp4wfp4_preshuffle_hot_exact_kernel[grid](
                A_q.view(torch.uint8),
                packed_w,
                y_pp,
                A_scale.view(torch.uint8),
                packed_scales,
                A_q.stride(0),
                A_q.stride(1),
                packed_w.stride(0),
                packed_w.stride(1),
                y_pp.stride(0),
                y_pp.stride(1),
                y_pp.stride(2),
                A_scale.stride(0),
                A_scale.stride(1),
                packed_scales.stride(0),
                packed_scales.stride(1),
                num_warps=local_config["num_warps"],
                num_stages=local_config["num_stages"],
                waves_per_eu=local_config["waves_per_eu"],
                matrix_instr_nonkdim=local_config["matrix_instr_nonkdim"],
                cache_modifier=local_config["cache_modifier"],
            )
    else:
        _vendored_gemm_afp4wfp4_preshuffle_hot_exact_kernel[grid](
            A_q.view(torch.uint8),
            packed_w,
            y_pp,
            A_scale.view(torch.uint8),
            packed_scales,
            A_q.stride(0),
            A_q.stride(1),
            packed_w.stride(0),
            packed_w.stride(1),
            y_pp.stride(0),
            y_pp.stride(1),
            y_pp.stride(2),
            A_scale.stride(0),
            A_scale.stride(1),
            packed_scales.stride(0),
            packed_scales.stride(1),
            num_warps=local_config["num_warps"],
            num_stages=local_config["num_stages"],
            waves_per_eu=local_config["waves_per_eu"],
            matrix_instr_nonkdim=local_config["matrix_instr_nonkdim"],
            cache_modifier=local_config["cache_modifier"],
        )

    grid_reduce = (33,)
    _vendored_gemm_afp4wfp4_reduce_hot_exact_kernel[grid_reduce](
        y_pp,
        y,
        y_pp.stride(0),
        y_pp.stride(1),
        y_pp.stride(2),
        y.stride(0),
        y.stride(1),
        BLOCK_SIZE_N=64,
    )
    return y


def custom_kernel(data: input_t) -> output_t:
    """
    Hybrid path:
    - exact `(16, 2112, 7168)` benchmark tuple uses the vendored preshuffled Triton path
    - everything else keeps the explicit asm or reference a4w4 fallback path
    """
    import aiter
    from aiter import dtypes

    A, _B, _B_q, B_shuffle, B_scale_sh = data
    B_shuffle = B_shuffle.contiguous()
    B_scale_sh = B_scale_sh.contiguous()
    m, k = A.shape
    n = B_shuffle.shape[0]
    shape = (m, n, k)

    if (
        shape in _FUSED_A16WFP4_SHAPES
        and triton is not None
        and _A16WFP4_PRESHUFFLE_STATE["enabled"]
    ):
        try:
            return _run_vendored_triton_a16wfp4_preshuffle_contract(A, B_shuffle, B_scale_sh)
        except Exception as exc:
            # One-shot disable: if the fused preshuffle kernel fails to compile or run in the
            # runner, keep the known-good paths for the remainder of the process.
            _A16WFP4_PRESHUFFLE_STATE["enabled"] = False
            _A16WFP4_PRESHUFFLE_STATE["error"] = repr(exc)

    vendored_config = _VENDORED_TRITON_PRESHUFFLE_CONFIGS.get(shape)
    if vendored_config is not None and triton is not None:
        # The preshuffled AFP4/WFP4 path requires unshuffled A scales for M < 32.
        if shape == (16, 2112, 7168):
            A_q, A_scale = _shape_tuned_quant_hot_exact(A)
            return _run_vendored_triton_preshuffle_hot_exact(
                A_q.contiguous(),
                B_shuffle,
                A_scale.contiguous(),
                B_scale_sh,
                vendored_config,
            )
        A_q, A_scale = _shape_tuned_quant(A, shuffle=False)
        return _run_vendored_triton_preshuffle(
            A_q.contiguous(),
            B_shuffle,
            A_scale.contiguous(),
            B_scale_sh,
            vendored_config,
        )

    A_q, A_scale_sh = _shape_tuned_quant(A, shuffle=True)

    kernel_name = _EXPLICIT_ASM_KERNELS.get(shape)
    if kernel_name is not None:
        log2_k_split = _EXPLICIT_ASM_LOG2_K_SPLIT.get(shape, _ASM_SPLIT_K)
        return _run_explicit_asm(
            aiter,
            A_q,
            B_shuffle,
            A_scale_sh,
            B_scale_sh,
            m,
            n,
            kernel_name,
            log2_k_split,
        )

    return aiter.gemm_a4w4(
        A_q,
        B_shuffle,
        A_scale_sh,
        B_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
scrolls · 1455 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