Skip to content
KernelIndex
Search⌘K

submission 574967

HorizonLiang · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

amd_202602_mxfp4_mm_large_self_gemm_64only.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-574967?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.5µs
#337 of 1143
2026-03-17

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:c668a9fe8ed0354066f9965a6328cc55eee600d7eb8fd2f42155b91f72ab54e3
license declaredunknown
license concludedunknown
authorsHorizonLiang
imported2026-08-26

Techniques

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

fp4Hybrid MXFP4 GEMM submission for MI355X.
num-warps = 4num_warps=4,
stages = 1num_stages = 1

Kernel source

amd_202602_mxfp4_mm_large_self_gemm_64only.py357 lines
"""
Hybrid MXFP4 GEMM submission for MI355X.

Small-M keeps the current strongest path. Large benchmark shapes reuse the fast
quantization path, then switch GEMM to Triton's AFP4xWFP4 preshuffled kernel so
we can isolate the custom GEMM ceiling without also changing quantization.
This variant only swaps the 64x7168x2048 benchmark shape.
"""

from __future__ import annotations

import triton
import triton.language as tl
import torch
from task import input_t, output_t
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op


@triton.heuristics(
    {
        "EVEN_M_N": lambda args: args["M"] % args["BLOCK_SIZE_M"] == 0
        and args["N"] % (args["BLOCK_SIZE_N"] * args["NUM_ITER"]) == 0,
    }
)
@triton.jit
def _dynamic_mxfp4_quant_kernel_shuffled(
    x_ptr,
    x_fp4_ptr,
    scale_sh_ptr,
    stride_x_m_in,
    stride_x_n_in,
    stride_x_fp4_m_in,
    stride_x_fp4_n_in,
    M,
    N,
    SCALE_N_VALID: tl.constexpr,
    SCALE_M_PAD: tl.constexpr,
    SCALE_N_PAD: tl.constexpr,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_N: tl.constexpr,
    NUM_ITER: tl.constexpr,
    NUM_STAGES: tl.constexpr,
    MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
    EVEN_M_N: tl.constexpr,
):
    pid_m = tl.program_id(0)
    start_n = tl.program_id(1) * NUM_ITER

    stride_x_m = tl.cast(stride_x_m_in, tl.int64)
    stride_x_n = tl.cast(stride_x_n_in, tl.int64)
    stride_x_fp4_m = tl.cast(stride_x_fp4_m_in, tl.int64)
    stride_x_fp4_n = tl.cast(stride_x_fp4_n_in, tl.int64)

    num_quant_blocks: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE

    for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):
        x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
        x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n

        if EVEN_M_N:
            x = tl.load(x_ptr + x_offs, cache_modifier=".cg").to(tl.float32)
        else:
            x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
            x = tl.load(x_ptr + x_offs, mask=x_mask, cache_modifier=".cg").to(
                tl.float32
            )

        x_fp4, scale_e8m0 = _mxfp4_quant_op(
            x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE
        )

        out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        out_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
        out_offs = (
            out_offs_m[:, None] * stride_x_fp4_m
            + out_offs_n[None, :] * stride_x_fp4_n
        )
        if EVEN_M_N:
            tl.store(x_fp4_ptr + out_offs, x_fp4)
        else:
            out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
            tl.store(x_fp4_ptr + out_offs, x_fp4, mask=out_mask)

        scale_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        scale_offs_n = pid_n * num_quant_blocks + tl.arange(0, num_quant_blocks)

        scale_offs_0 = scale_offs_m[:, None] // 32
        scale_offs_1 = scale_offs_m[:, None] % 32
        scale_offs_2 = scale_offs_1 % 16
        scale_offs_1 = scale_offs_1 // 16
        scale_offs_3 = scale_offs_n[None, :] // 8
        scale_offs_4 = scale_offs_n[None, :] % 8
        scale_offs_5 = scale_offs_4 % 4
        scale_offs_4 = scale_offs_4 // 4
        scale_offs = (
            scale_offs_1
            + scale_offs_4 * 2
            + scale_offs_2 * 4
            + scale_offs_5 * 64
            + scale_offs_3 * 256
            + scale_offs_0 * 32 * SCALE_N_VALID
        )

        scale_mask_valid = (scale_offs_m < M)[:, None] & (
            scale_offs_n < ((N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE)
        )[None, :]
        scale_mask_pad = (scale_offs_m < SCALE_M_PAD)[:, None] & (
            scale_offs_n < SCALE_N_PAD
        )[None, :]
        scale_e8m0 = tl.where(scale_mask_valid, scale_e8m0, 127)
        tl.store(scale_sh_ptr + scale_offs, scale_e8m0, mask=scale_mask_pad)


_QUANT_CACHE: dict[tuple[tuple[str, int | None], int, int], tuple[torch.Tensor, torch.Tensor]] = {}
_OUTPUT_CACHE: dict[tuple[tuple[str, int | None], int, int], torch.Tensor] = {}

_KERNEL_OVERRIDES: dict[tuple[int, int, int], str] = {
    (4, 2880, 512): "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E",
    (16, 2112, 7168): "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
    (32, 4096, 512): "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E",
    (32, 2880, 512): "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E",
    (64, 7168, 2048): "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
    (256, 3072, 1536): "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E",
}

_SMALL_M_CONFIG_OVERRIDES: dict[tuple[int, int, int], dict[str, int | str | None]] = {
    (4, 2880, 512): {
        "BLOCK_SIZE_M": 4,
        "BLOCK_SIZE_N": 64,
        "BLOCK_SIZE_K": 512,
        "GROUP_SIZE_M": 1,
        "num_warps": 2,
        "num_stages": 1,
        "waves_per_eu": 2,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": None,
        "NUM_KSPLIT": 1,
    },
    (16, 2112, 7168): {
        "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": 7,
    },
}

_LARGE_SELF_GEMM_CONFIGS: dict[tuple[int, int, int], dict[str, int | str | None]] = {
    (64, 7168, 2048): {
        "BLOCK_SIZE_M": 32,
        "BLOCK_SIZE_N": 16,
        "BLOCK_SIZE_K": 1024,
        "GROUP_SIZE_M": 1,
        "num_warps": 2,
        "num_stages": 2,
        "waves_per_eu": 2,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": None,
        "NUM_KSPLIT": 1,
    },
}


def _device_key(device: torch.device) -> tuple[str, int | None]:
    return device.type, device.index


def _get_quant_outputs(
    m: int,
    n: int,
    device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor]:
    scale_n_valid = triton.cdiv(n, 32)
    scale_n_pad = triton.cdiv(scale_n_valid, 8) * 8
    scale_m_pad = triton.cdiv(m, 256) * 256
    key = (_device_key(device), m, n)
    cached = _QUANT_CACHE.get(key)
    if cached is None:
        cached = (
            torch.empty((m, n // 2), dtype=torch.uint8, device=device),
            torch.empty((scale_m_pad, scale_n_pad), dtype=torch.uint8, device=device),
        )
        _QUANT_CACHE[key] = cached
    return cached


def _get_output(m: int, n: int, device: torch.device) -> torch.Tensor:
    key = (_device_key(device), m, n)
    cached = _OUTPUT_CACHE.get(key)
    if cached is None:
        cached = torch.empty((m, n), dtype=torch.bfloat16, device=device)
        _OUTPUT_CACHE[key] = cached
    return cached


def _quant_mxfp4_shuffled(x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    m, n = x.shape
    x_fp4, scale_sh = _get_quant_outputs(m, n, x.device)

    if m <= 32:
        num_iter = 1
        block_size_m = triton.next_power_of_2(m)
        block_size_n = 32
        num_stages = 1
    else:
        num_iter = 4
        block_size_m = 64
        block_size_n = 64
        num_stages = 2
        if n <= 16384:
            block_size_m = 32
            block_size_n = 128

    if n <= 1024:
        num_iter = 1
        num_stages = 1
        block_size_n = max(32, min(256, triton.next_power_of_2(n)))
        block_size_m = min(8, triton.next_power_of_2(m))

    grid = (
        triton.cdiv(m, block_size_m),
        triton.cdiv(n, block_size_n * num_iter),
    )
    _dynamic_mxfp4_quant_kernel_shuffled[grid](
        x,
        x_fp4,
        scale_sh,
        *x.stride(),
        *x_fp4.stride(),
        M=m,
        N=n,
        SCALE_N_VALID=triton.cdiv(n, 32),
        SCALE_M_PAD=triton.cdiv(m, 256) * 256,
        SCALE_N_PAD=triton.cdiv(triton.cdiv(n, 32), 8) * 8,
        BLOCK_SIZE_M=block_size_m,
        BLOCK_SIZE_N=block_size_n,
        NUM_ITER=num_iter,
        NUM_STAGES=num_stages,
        MXFP4_QUANT_BLOCK_SIZE=32,
        num_warps=4,
        num_stages=1,
    )
    return x_fp4, scale_sh


def _view_preshuffled_weight(
    weight_shuffle: torch.Tensor,
    weight_scale_shuffle: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
    n, k_half = weight_shuffle.shape
    weight_u8 = (
        weight_shuffle.view(torch.uint8)
        if weight_shuffle.dtype is not torch.uint8
        else weight_shuffle
    )
    scale_u8 = (
        weight_scale_shuffle.view(torch.uint8)
        if weight_scale_shuffle.dtype is not torch.uint8
        else weight_scale_shuffle
    )
    scale_m, scale_n = scale_u8.shape
    weight_triton = weight_u8.view(n // 16, k_half * 16)
    scale_triton = scale_u8.view(scale_m // 32, scale_n * 32)[: n // 32]
    return weight_triton, scale_triton


def _view_preshuffled_activation_scales(
    scale_shuffle: torch.Tensor,
    m: int,
) -> torch.Tensor:
    scale_u8 = (
        scale_shuffle.view(torch.uint8)
        if scale_shuffle.dtype is not torch.uint8
        else scale_shuffle
    )
    scale_m, scale_n = scale_u8.shape
    return scale_u8.view(scale_m // 32, scale_n * 32)[: m // 32]


def custom_kernel(data: input_t) -> output_t:
    import aiter
    from aiter import dtypes
    from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
    from aiter.ops.triton.gemm.basic.gemm_a16wfp4 import gemm_a16wfp4_preshuffle
    from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import gemm_afp4wfp4_preshuffle

    a, b, b_q, b_shuffle, b_scale_sh = data
    del b, b_q

    a = a.contiguous()
    m, k = a.shape
    n = b_shuffle.shape[0]
    shape = (m, n, k)

    if m <= 32:
        w_triton, w_scale_triton = _view_preshuffled_weight(b_shuffle, b_scale_sh)
        out = _get_output(m, n, a.device)
        config = _SMALL_M_CONFIG_OVERRIDES.get(shape)
        return gemm_a16wfp4_preshuffle(
            a,
            w_triton,
            w_scale_triton,
            prequant=True,
            dtype=torch.bfloat16,
            y=out,
            config=config,
        )

    a_q, a_scale_sh = _quant_mxfp4_shuffled(a)
    large_config = _LARGE_SELF_GEMM_CONFIGS.get(shape)
    if large_config is not None:
        w_triton, w_scale_triton = _view_preshuffled_weight(b_shuffle, b_scale_sh)
        a_scale_triton = _view_preshuffled_activation_scales(a_scale_sh, m)
        out = _get_output(m, n, a.device)
        return gemm_afp4wfp4_preshuffle(
            a_q,
            w_triton,
            a_scale_triton,
            w_scale_triton,
            dtype=torch.bfloat16,
            y=out,
            config=large_config,
        )

    kernel_name = _KERNEL_OVERRIDES.get(shape)
    if kernel_name is None:
        return aiter.gemm_a4w4(
            a_q.view(dtypes.fp4x2),
            b_shuffle,
            a_scale_sh.view(dtypes.fp8_e8m0),
            b_scale_sh,
            dtype=dtypes.bf16,
            bpreshuffle=True,
        )

    out = _get_output(m, n, a.device)
    gemm_a4w4_asm(
        a_q.view(dtypes.fp4x2),
        b_shuffle,
        a_scale_sh.view(dtypes.fp8_e8m0),
        b_scale_sh,
        out,
        kernel_name,
        None,
        1.0,
        0.0,
        True,
        0,
    )
    return out
scrolls · 357 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