Skip to content
KernelIndex
Search⌘K

submission 746650

guangxiangdebizi · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_exp17_v05.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-746650?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.0µs
#303 of 1143
2026-04-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f016da420ce7d4ef12098fc742e9884b675934dade3c65c0900a1551f398fed3
license declaredunknown
license concludedunknown
authorsguangxiangdebizi
imported2026-08-26

Techniques

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

fp4"""Experiment 17 v05 for AMD MXFP4 GEMM.
split-kand (args["SPLITK_BLOCK_SIZE"] % args["BLOCK_SIZE_K"] == 0)
tile-k = 512- Shape 5 (64,7168,2048): keep fused BM=16 BN=128 BK=512 (BN=256 regressed)
tile-m = 16- Shape 5 (64,7168,2048): keep fused BM=16 BN=128 BK=512 (BN=256 regressed)
tile-n = 128- Shape 5 (64,7168,2048): keep fused BM=16 BN=128 BK=512 (BN=256 regressed)

Kernel source

submission_exp17_v05.py988 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X

"""Experiment 17 v05 for AMD MXFP4 GEMM.

Cherry-pick best configs on top of exp15_v01:
- Shape 2 (16,2112,7168): num_stages 1->2 (validated -0.8µs in exp17_v04)
- Shape 5 (64,7168,2048): keep fused BM=16 BN=128 BK=512 (BN=256 regressed)
- Shape 6 (256,3072,1536): keep two-stage
"""

from __future__ import annotations

import torch

from task import input_t, output_t

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


def _cfg(
    block_size_m: int,
    block_size_n: int,
    block_size_k: int,
    group_size_m: int,
    num_warps: int,
    num_stages: int,
    waves_per_eu: int,
    matrix_instr_nonkdim: int,
    cache_modifier: str | None,
    num_ksplit: int,
) -> dict[str, int | str | None]:
    return {
        "BLOCK_SIZE_M": block_size_m,
        "BLOCK_SIZE_N": block_size_n,
        "BLOCK_SIZE_K": block_size_k,
        "GROUP_SIZE_M": group_size_m,
        "num_warps": num_warps,
        "num_stages": num_stages,
        "waves_per_eu": waves_per_eu,
        "matrix_instr_nonkdim": matrix_instr_nonkdim,
        "cache_modifier": cache_modifier,
        "NUM_KSPLIT": num_ksplit,
    }


_DEFAULT_CONFIG = _cfg(32, 64, 512, 1, 8, 1, 2, 16, None, 1)
_SMALL_K_TINY_M_CONFIG = _cfg(4, 128, 512, 1, 4, 1, 2, 16, ".cg", 1)
_SMALL_K_SMALL_M_CONFIG = _cfg(16, 128, 512, 1, 4, 1, 2, 16, None, 1)
_SMALL_K_4096_CONFIG = _cfg(8, 128, 512, 1, 8, 1, 2, 16, ".cg", 1)
_LARGE_M_CONFIG = _cfg(32, 256, 256, 4, 8, 2, 4, 16, None, 1)
_TWO_STAGE_64_CONFIG = _cfg(16, 128, 512, 1, 8, 2, 4, 16, None, 1)
_TWO_STAGE_256_CONFIG = _cfg(16, 256, 256, 4, 8, 2, 4, 16, None, 1)

_SPECIALIZED_NK_CONFIGS: dict[
    tuple[int, int], list[tuple[int | None, dict[str, int | str | None]]]
] = {
    (2112, 7168): [
        (8, _cfg(8, 128, 512, 1, 4, 1, 1, 16, ".cg", 14)),
        (16, _cfg(16, 128, 512, 1, 4, 1, 1, 16, ".cg", 14)),
        (32, _cfg(16, 128, 512, 1, 4, 1, 2, 16, None, 14)),
        (64, _cfg(16, 128, 512, 1, 4, 1, 2, 16, None, 14)),
        (128, _cfg(32, 128, 512, 1, 4, 1, 2, 16, None, 14)),
        (256, _cfg(32, 128, 512, 1, 4, 1, 2, 16, None, 14)),
        (None, _cfg(32, 128, 256, 4, 2, 2, 2, 16, None, 1)),
    ],
    (7168, 2048): [
        (8, _cfg(8, 128, 512, 1, 8, 2, 1, 16, ".cg", 4)),
        (16, _cfg(16, 128, 512, 1, 4, 2, 2, 16, ".cg", 4)),
        (32, _cfg(16, 128, 512, 1, 8, 2, 2, 16, ".cg", 4)),
        (64, _cfg(16, 128, 512, 1, 8, 2, 4, 16, None, 1)),
        (128, _cfg(32, 128, 256, 4, 8, 2, 4, 16, None, 1)),
        (256, _cfg(32, 256, 256, 4, 8, 2, 4, 16, None, 1)),
        (None, _cfg(32, 256, 256, 1, 8, 2, 1, 16, None, 1)),
    ],
    (3072, 1536): [
        (16, _cfg(16, 64, 256, 1, 4, 2, 2, 16, None, 4)),
        (64, _cfg(32, 256, 256, 4, 8, 2, 4, 16, None, 1)),
        (256, _cfg(32, 256, 256, 4, 8, 2, 4, 16, None, 1)),
        (None, _cfg(32, 256, 256, 4, 8, 2, 4, 16, None, 1)),
    ],
    (4096, 512): [
        (32, _SMALL_K_4096_CONFIG),
        (None, _DEFAULT_CONFIG),
    ],
}

_FIXED_SHAPE_CONFIGS: dict[tuple[int, int, int], dict[str, int | str | None]] = {
    (4, 2880, 512): _SMALL_K_TINY_M_CONFIG,
    (16, 2112, 7168): _cfg(16, 128, 512, 1, 4, 2, 1, 16, ".cg", 14),
    (32, 4096, 512): _SMALL_K_4096_CONFIG,
    (32, 2880, 512): _SMALL_K_SMALL_M_CONFIG,
    (64, 7168, 2048): _cfg(16, 128, 512, 1, 8, 2, 4, 16, None, 1),
    (256, 3072, 1536): _LARGE_M_CONFIG,
}

_TWO_STAGE_FIXED_SHAPE_CONFIGS: dict[tuple[int, int, int], dict[str, int | str | None]] = {
    (256, 3072, 1536): _TWO_STAGE_256_CONFIG,
}


if triton is not None:

    @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.jit
    def _mxfp4_quant_op(
        x,
        BLOCK_SIZE_N: tl.constexpr,
        BLOCK_SIZE_M: tl.constexpr,
        MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
    ):
        exp_bias_fp32: tl.constexpr = 127
        exp_bias_fp4: tl.constexpr = 1
        ebits_fp32: tl.constexpr = 8
        ebits_fp4: tl.constexpr = 2
        mbits_fp32: tl.constexpr = 23
        mbits_fp4: tl.constexpr = 1

        max_normal: tl.constexpr = 6
        min_normal: tl.constexpr = 1

        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)

        bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127
        quant_scale = tl.exp2(-scale_e8m0_unbiased)

        qx = (x * quant_scale).to(tl.uint32, bitcast=True)
        sign = qx & 0x80000000
        qx = qx ^ sign

        qx_fp32 = qx.to(tl.float32, bitcast=True)
        saturate_mask = qx_fp32 >= max_normal
        denormal_mask = (~saturate_mask) & (qx_fp32 < min_normal)
        normal_mask = ~(saturate_mask | denormal_mask)

        denorm_exp: tl.constexpr = (
            (exp_bias_fp32 - exp_bias_fp4) + (mbits_fp32 - mbits_fp4) + 1
        )
        denorm_mask_int: tl.constexpr = denorm_exp << mbits_fp32
        denorm_mask_float: tl.constexpr = tl.cast(
            denorm_mask_int, tl.float32, bitcast=True
        )

        denormal_x = qx_fp32 + denorm_mask_float
        denormal_x = denormal_x.to(tl.uint32, bitcast=True)
        denormal_x -= denorm_mask_int
        denormal_x = denormal_x.to(tl.uint8)

        normal_x = qx
        mant_odd = (normal_x >> (mbits_fp32 - mbits_fp4)) & 1
        val_to_add = ((exp_bias_fp4 - exp_bias_fp32) << mbits_fp32) + (1 << 21) - 1
        normal_x += val_to_add
        normal_x += mant_odd
        normal_x = (normal_x >> (mbits_fp32 - mbits_fp4)).to(tl.uint8)

        e2m1_value = tl.full(qx.type.get_block_shapes(), 0x7, dtype=tl.uint8)
        e2m1_value = tl.where(normal_mask, normal_x, e2m1_value)
        e2m1_value = tl.where(denormal_mask, denormal_x, e2m1_value)

        sign_lp = sign >> (mbits_fp32 + ebits_fp32 - mbits_fp4 - ebits_fp4)
        e2m1_value = e2m1_value | sign_lp.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)).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"] == 0})
    @triton.jit
    def _mxfp4_quant_matrix_kernel(
        a_ptr,
        a_fp4_ptr,
        a_scales_ptr,
        M,
        K,
        stride_am,
        stride_ak,
        stride_afp4_m,
        stride_afp4_k,
        stride_asm,
        stride_ask,
        BLOCK_SIZE_M: tl.constexpr,
        BLOCK_SIZE_K: tl.constexpr,
        EVEN_K: tl.constexpr,
    ):
        pid_m = tl.program_id(axis=0)
        pid_k = tl.program_id(axis=1)

        offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        offs_k = pid_k * BLOCK_SIZE_K + tl.arange(0, BLOCK_SIZE_K)
        a_ptrs = a_ptr + offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak

        if EVEN_K:
            a_bf16 = tl.load(a_ptrs)
        else:
            a_bf16 = tl.load(a_ptrs, mask=(offs_m[:, None] < M) & (offs_k[None, :] < K), other=0)

        a_fp4, a_scales = _mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)

        offs_k_fp4 = pid_k * (BLOCK_SIZE_K // 2) + tl.arange(0, BLOCK_SIZE_K // 2)
        a_fp4_ptrs = (
            a_fp4_ptr
            + offs_m[:, None] * stride_afp4_m
            + offs_k_fp4[None, :] * stride_afp4_k
        )
        tl.store(a_fp4_ptrs, a_fp4, mask=offs_m[:, None] < M)

        offs_k_scale = pid_k * (BLOCK_SIZE_K // 32) + tl.arange(0, BLOCK_SIZE_K // 32)
        a_scale_ptrs = (
            a_scales_ptr
            + offs_m[:, None] * stride_asm
            + offs_k_scale[None, :] * stride_ask
        )
        tl.store(a_scale_ptrs, a_scales, mask=offs_m[:, None] < M)


    @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 _gemm_a16wfp4_prequant_kernel(
        a_ptr,
        a_scales_ptr,
        b_ptr,
        c_ptr,
        b_scales_ptr,
        M,
        N,
        K,
        stride_am,
        stride_ak,
        stride_asm,
        stride_ask,
        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_asm > 0)
        tl.assume(stride_ask > 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
        num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
        num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)

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

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

        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_fp4 = tl.arange(0, BLOCK_SIZE_K // 2)
            offs_k_split_fp4 = pid_k * (SPLITK_BLOCK_SIZE // 2) + offs_k_fp4
            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_fp4[None, :] * stride_ak
            )

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

            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_b = (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_b[None, :] * stride_bsk
            )

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

            for k_iter 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_fp4[None, :] < K - k_iter * (BLOCK_SIZE_K // 2),
                        other=0,
                    )
                    b = tl.load(
                        b_ptrs,
                        mask=offs_k_shuffle_arr[None, :]
                        < (K - k_iter * (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")

                a_ptrs += (BLOCK_SIZE_K // 2) * stride_ak
                a_scale_ptrs += (BLOCK_SIZE_K // scale_group_size) * stride_ask
                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)


    @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 _gemm_a16wfp4_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,
        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
        num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
        num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)

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

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

        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_iter 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)
                    b = tl.load(b_ptrs, cache_modifier=cache_modifier)
                else:
                    a_bf16 = tl.load(
                        a_ptrs,
                        mask=offs_k_bf16[None, :] < 2 * K - k_iter * BLOCK_SIZE_K,
                        other=0,
                    )
                    b = tl.load(
                        b_ptrs,
                        mask=offs_k_shuffle_arr[None, :]
                        < (K - k_iter * (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)
                )

                a, a_scales = _mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
                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)


    @triton.jit
    def _gemm_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).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)


def _pick_specialized_nk_config(
    m: int, n: int, k: int
) -> dict[str, int | str | None] | None:
    configs = _SPECIALIZED_NK_CONFIGS.get((n, k))
    if configs is None:
        return None
    for upper_bound, config in configs:
        if upper_bound is None or m <= upper_bound:
            return config
    return None


def _pick_two_stage_config(
    m: int, n: int, k: int
) -> dict[str, int | str | None] | None:
    config = _TWO_STAGE_FIXED_SHAPE_CONFIGS.get((m, n, k))
    if config is None:
        return None
    return dict(config)


def _pick_config(m: int, n: int, k: int) -> dict[str, int | str | None]:
    exact = _FIXED_SHAPE_CONFIGS.get((m, n, k))
    if exact is not None:
        return dict(exact)

    specialized = _pick_specialized_nk_config(m, n, k)
    if specialized is not None:
        return dict(specialized)

    if k == 512 and m <= 8:
        return dict(_SMALL_K_TINY_M_CONFIG)
    if m >= 128 and n >= 2048 and k >= 1024:
        return dict(_LARGE_M_CONFIG)
    return dict(_DEFAULT_CONFIG)


def _get_splitk(k_packed: int, block_size_k: int, num_ksplit: int) -> tuple[int, int, int]:
    if triton is None:
        raise RuntimeError("Triton is required for the fused MXFP4 kernel.")

    splitk_block_size = (
        triton.cdiv((2 * triton.cdiv(k_packed, num_ksplit)), block_size_k)
        * block_size_k
    )
    while num_ksplit > 1 and block_size_k > 16:
        if (
            k_packed % (splitk_block_size // 2) == 0
            and splitk_block_size % block_size_k == 0
            and k_packed % (block_size_k // 2) == 0
        ):
            break
        if k_packed % (splitk_block_size // 2) != 0 and num_ksplit > 1:
            num_ksplit //= 2
        elif splitk_block_size % block_size_k != 0:
            if num_ksplit > 1:
                num_ksplit //= 2
            elif block_size_k > 16:
                block_size_k //= 2
        elif k_packed % (block_size_k // 2) != 0 and block_size_k > 16:
            block_size_k //= 2
        else:
            break
        splitk_block_size = (
            triton.cdiv((2 * triton.cdiv(k_packed, num_ksplit)), block_size_k)
            * block_size_k
        )
    return splitk_block_size, block_size_k, num_ksplit


def _reshape_b_shuffle_for_preshuffle(b_shuffle: torch.Tensor) -> torch.Tensor:
    n, k_half = b_shuffle.shape
    if n % 16 != 0:
        raise ValueError(f"Expected N to be divisible by 16, but got N={n}.")
    b_in = b_shuffle if b_shuffle.is_contiguous() else b_shuffle.contiguous()
    b_bytes = b_in.view(torch.uint8)
    return b_bytes.view(n // 16, k_half * 16)


def _reshape_b_scale_for_preshuffle(
    b_scale_sh: torch.Tensor,
    n: int,
    k_bf16: int,
) -> torch.Tensor:
    if n % 32 != 0:
        raise ValueError(f"Expected N to be divisible by 32, but got N={n}.")
    k_scale = k_bf16 // 32
    scale_slice = b_scale_sh[:n, :k_scale]
    scale_in = scale_slice if scale_slice.is_contiguous() else scale_slice.contiguous()
    scale_bytes = scale_in.view(torch.uint8)
    return scale_bytes.view(n // 32, k_scale * 32)


def _gemm_a16wfp4_preshuffle(
    a: torch.Tensor,
    b_shuffle: torch.Tensor,
    b_scale_sh: torch.Tensor,
    dtype: torch.dtype = torch.bfloat16,
) -> torch.Tensor:
    if triton is None:
        raise RuntimeError("Triton is not available in this environment.")

    m, k_bf16 = a.shape
    n_blocks, packed_k_x16 = b_shuffle.shape
    if packed_k_x16 % 16 != 0:
        raise ValueError(
            f"Expected preshuffled B second dim to be divisible by 16, got {packed_k_x16}."
        )

    n = n_blocks * 16
    k_packed = packed_k_x16 // 16
    if 2 * k_packed != k_bf16:
        raise ValueError(
            "Unexpected preshuffled B shape: "
            f"A has bf16 K={k_bf16}, but B encodes packed K={k_packed}."
        )

    config = _pick_config(m, n, k_bf16)
    num_ksplit = int(config["NUM_KSPLIT"])
    block_size_k = int(config["BLOCK_SIZE_K"])

    if num_ksplit > 1:
        splitk_block_size, block_size_k, num_ksplit = _get_splitk(
            k_packed, block_size_k, num_ksplit
        )
        config["SPLITK_BLOCK_SIZE"] = splitk_block_size
        config["BLOCK_SIZE_K"] = block_size_k
        config["NUM_KSPLIT"] = num_ksplit

    if int(config["BLOCK_SIZE_K"]) >= 2 * k_packed:
        config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * k_packed)
        config["SPLITK_BLOCK_SIZE"] = 2 * k_packed
        config["NUM_KSPLIT"] = 1
    else:
        config["SPLITK_BLOCK_SIZE"] = (
            int(config["SPLITK_BLOCK_SIZE"])
            if "SPLITK_BLOCK_SIZE" in config
            else 2 * k_packed
        )

    config["BLOCK_SIZE_N"] = max(int(config["BLOCK_SIZE_N"]), 32)

    if int(config["NUM_KSPLIT"]) > 1:
        y_pp = torch.empty(
            (int(config["NUM_KSPLIT"]), m, n),
            dtype=torch.float32,
            device=a.device,
        )
        y = torch.empty((m, n), dtype=dtype, device=a.device)
    else:
        config["SPLITK_BLOCK_SIZE"] = 2 * k_packed
        y_pp = None
        y = torch.empty((m, n), dtype=dtype, device=a.device)

    grid = lambda meta: (  # noqa: E731
        int(meta["NUM_KSPLIT"])
        * triton.cdiv(m, int(meta["BLOCK_SIZE_M"]))
        * triton.cdiv(n, int(meta["BLOCK_SIZE_N"])),
    )

    _gemm_a16wfp4_preshuffle_kernel[grid](
        a,
        b_shuffle,
        y if y_pp is None else y_pp,
        b_scale_sh,
        m,
        n,
        k_packed,
        a.stride(0),
        a.stride(1),
        b_shuffle.stride(0),
        b_shuffle.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),
        b_scale_sh.stride(0),
        b_scale_sh.stride(1),
        **config,
    )

    if y_pp is not None:
        reduce_block_m = 16
        reduce_block_n = 64
        actual_ksplit = triton.cdiv(k_packed, int(config["SPLITK_BLOCK_SIZE"]) // 2)
        grid_reduce = (
            triton.cdiv(m, reduce_block_m),
            triton.cdiv(n, reduce_block_n),
        )
        _gemm_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),
            BLOCK_SIZE_M=reduce_block_m,
            BLOCK_SIZE_N=reduce_block_n,
            ACTUAL_KSPLIT=actual_ksplit,
            MAX_KSPLIT=triton.next_power_of_2(int(config["NUM_KSPLIT"])),
        )

    return y


def _gemm_a16wfp4_two_stage(
    a: torch.Tensor,
    b_shuffle: torch.Tensor,
    b_scale_sh: torch.Tensor,
    dtype: torch.dtype = torch.bfloat16,
) -> torch.Tensor:
    if triton is None:
        raise RuntimeError("Triton is not available in this environment.")

    m, k_bf16 = a.shape
    n_blocks, packed_k_x16 = b_shuffle.shape
    if packed_k_x16 % 16 != 0:
        raise ValueError(
            f"Expected preshuffled B second dim to be divisible by 16, got {packed_k_x16}."
        )

    n = n_blocks * 16
    k_packed = packed_k_x16 // 16
    if 2 * k_packed != k_bf16:
        raise ValueError(
            "Unexpected preshuffled B shape: "
            f"A has bf16 K={k_bf16}, but B encodes packed K={k_packed}."
        )

    config = _pick_two_stage_config(m, n, k_bf16)
    if config is None:
        raise ValueError(f"Two-stage config not defined for shape {(m, n, k_bf16)}.")

    if int(config["NUM_KSPLIT"]) != 1:
        raise ValueError("Two-stage prototype currently expects NUM_KSPLIT == 1.")

    block_size_k = int(config["BLOCK_SIZE_K"])
    config["BLOCK_SIZE_N"] = max(int(config["BLOCK_SIZE_N"]), 32)

    a_fp4 = torch.empty((m, k_bf16 // 2), dtype=torch.uint8, device=a.device)
    a_scales = torch.empty((m, k_bf16 // 32), dtype=torch.uint8, device=a.device)
    y = torch.empty((m, n), dtype=dtype, device=a.device)

    grid_quant = (
        triton.cdiv(m, int(config["BLOCK_SIZE_M"])),
        triton.cdiv(k_bf16, block_size_k),
    )
    _mxfp4_quant_matrix_kernel[grid_quant](
        a,
        a_fp4,
        a_scales,
        m,
        k_bf16,
        a.stride(0),
        a.stride(1),
        a_fp4.stride(0),
        a_fp4.stride(1),
        a_scales.stride(0),
        a_scales.stride(1),
        BLOCK_SIZE_M=int(config["BLOCK_SIZE_M"]),
        BLOCK_SIZE_K=block_size_k,
    )

    grid = lambda meta: (  # noqa: E731
        triton.cdiv(m, int(meta["BLOCK_SIZE_M"]))
        * triton.cdiv(n, int(meta["BLOCK_SIZE_N"])),
    )

    _gemm_a16wfp4_prequant_kernel[grid](
        a_fp4,
        a_scales,
        b_shuffle,
        y,
        b_scale_sh,
        m,
        n,
        k_packed,
        a_fp4.stride(0),
        a_fp4.stride(1),
        a_scales.stride(0),
        a_scales.stride(1),
        b_shuffle.stride(0),
        b_shuffle.stride(1),
        0,
        y.stride(0),
        y.stride(1),
        b_scale_sh.stride(0),
        b_scale_sh.stride(1),
        BLOCK_SIZE_M=int(config["BLOCK_SIZE_M"]),
        BLOCK_SIZE_N=int(config["BLOCK_SIZE_N"]),
        BLOCK_SIZE_K=block_size_k,
        GROUP_SIZE_M=int(config["GROUP_SIZE_M"]),
        NUM_KSPLIT=1,
        SPLITK_BLOCK_SIZE=2 * k_packed,
        num_warps=int(config["num_warps"]),
        num_stages=int(config["num_stages"]),
        waves_per_eu=int(config["waves_per_eu"]),
        matrix_instr_nonkdim=int(config["matrix_instr_nonkdim"]),
        cache_modifier=config["cache_modifier"],
    )

    return y


@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
    a, _, _, b_shuffle, b_scale_sh = data
    a_in = a if a.is_contiguous() else a.contiguous()
    n = b_shuffle.shape[0]
    k_bf16 = a_in.shape[1]
    b_preshuffle = _reshape_b_shuffle_for_preshuffle(b_shuffle)
    b_scale_preshuffle = _reshape_b_scale_for_preshuffle(b_scale_sh, n, k_bf16)
    # Only Shape 6 keeps the extra A prequant pass in exp15.
    if (a_in.shape[0], n, k_bf16) in _TWO_STAGE_FIXED_SHAPE_CONFIGS:
        return _gemm_a16wfp4_two_stage(a_in, b_preshuffle, b_scale_preshuffle)
    return _gemm_a16wfp4_preshuffle(a_in, b_preshuffle, b_scale_preshuffle)
scrolls · 988 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