Skip to content
KernelIndex
Search⌘K

submission 721672

FelliYang · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

v53.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-721672?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
13.1µs
#408 of 1143
2026-04-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:40392b2bddb94e4397a6eb1552b252e138576134616d593a76b9f9be5dd81c28
license declaredunknown
license concludedunknown
authorsFelliYang
imported2026-08-26

Techniques

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

fp4MXFP4 GEMM - v53: v15_merge_v30 + Triton splitK for m=16, n=2112, k=7168.
num-warps = 4num_warps=4,
split-kMXFP4 GEMM - v53: v15_merge_v30 + Triton splitK for m=16, n=2112, k=7168.
stages = 1NUM_STAGES=NUM_STAGES, num_warps=NUM_WARPS, waves_per_eu=0, num_stages=1,

Kernel source

v53.py318 lines
"""
MXFP4 GEMM - v53: v15_merge_v30 + Triton splitK for m=16, n=2112, k=7168.

m=16, k=7168 原来: 17 CTAs (32x128 tile, no splitK) → 6.6% CU利用率
v53: NUM_KSPLIT=14 → 17×14=238 CTAs → ~93% CU利用率

关键变化:
  1. m=16 新增 _mxfp4_quant_natural_kernel: 输出 natural (M, K//32) A_scale
     (gemm_afp4wfp4_preshuffle 在 M<32 时期望 un-shuffled A_scale)
  2. B_scale_sh 直接 view 成 (sm//32, K): ASM shuffle format 与
     _shuffle_scales 输出在 K%256==0 时完全等价, 零拷贝
  3. B 相关 tensor 用全局变量而非 dict cache (B 是固定权重)
  4. 其余 shape 完全沿用 v15_merge_v30 的 ASM 路径
"""
import torch
import triton
import triton.language as tl

try:
    from task import input_t, output_t
except ImportError:
    from typing import Any, Tuple
    input_t = Tuple[Any, ...]
    output_t = Any

from aiter import dtypes
import aiter
from aiter.ops.triton.quant import _mxfp4_quant_op
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import gemm_afp4wfp4_preshuffle


# ── ASM path: small-M tile configs (m=4/32) ────────────────────────────────
SHAPE_CONFIGS = {
    (4,  2880, 512):  ("32x128", 0),
    (32, 4096, 512):  ("32x128", 0),
    (32, 2880, 512):  ("32x128", 0),
}

def _make_kernel_name(suffix):
    base = f"f4gemm_bf16_per1x32Fp4_BpreShuffle_{suffix}"
    return f"_ZN5aiter{len(base)}{base}E"


# ── Quant kernel for ASM path: shuffled A_scale ─────────────────────────────
@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 _fused_mxfp4_quant_shuffle_kernel(
    x_ptr, x_fp4_ptr, bs_ptr,
    stride_x_m_in, stride_x_n_in, stride_x_fp4_m_in, stride_x_fp4_n_in,
    M, N, scaleN: tl.int64, scaleN_pad: tl.int64,
    BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
    NUM_ITER: tl.constexpr, NUM_STAGES: tl.constexpr,
    MXFP4_QUANT_BLOCK_SIZE: tl.constexpr, EVEN_M_N: tl.constexpr,
    SCALING_MODE: tl.constexpr,
):
    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)

        out_tensor, bs_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, out_tensor)
        else:
            out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
            tl.store(x_fp4_ptr + out_offs, out_tensor, mask=out_mask)

        bs_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        bs_offs_n = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
        m_idx = bs_offs_m[:, None]
        n_idx = bs_offs_n[None, :]
        d0 = m_idx // 32; d5 = (m_idx % 32) // 16; d3 = m_idx % 16
        d1 = n_idx // 8;  d4 = (n_idx % 8) // 4;   d2 = n_idx % 4
        shuffle_offs = d0 * 32 * scaleN_pad + d1 * 256 + d2 * 64 + d3 * 4 + d4 * 2 + d5
        if EVEN_M_N:
            tl.store(bs_ptr + shuffle_offs, bs_e8m0)
        else:
            bs_mask = (bs_offs_m < M)[:, None] & (bs_offs_n < scaleN)[None, :]
            tl.store(bs_ptr + shuffle_offs, bs_e8m0, mask=bs_mask)


# ── Quant kernel for Triton/splitK path: natural (M, K//32) A_scale ─────────
# gemm_afp4wfp4_preshuffle 在 M<32 时期望 un-shuffled A_scale:
#   x_scale shape (M, K//32), stride (K//32, 1)
# Grid: (cdiv(M,BSM), cdiv(K,BSK)) — 对 m=16,k=7168 就是 (1, 14)
@triton.jit
def _mxfp4_quant_natural_kernel(
    x_ptr, fp4_ptr, scale_ptr,
    stride_xm, stride_xk,
    stride_fm, stride_fk,
    stride_sm, stride_sk,
    M, K,
    BLOCK_SIZE_M: tl.constexpr,
    BLOCK_SIZE_K: tl.constexpr,
    MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_k = tl.program_id(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)

    x = tl.load(
        x_ptr + offs_m[:, None] * stride_xm + offs_k[None, :] * stride_xk,
        cache_modifier=".cg",
    ).to(tl.float32)

    fp4_out, scale_out = _mxfp4_quant_op(x, BLOCK_SIZE_K, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE)

    # fp4: (BLOCK_SIZE_M, BLOCK_SIZE_K//2)
    offs_fk = pid_k * (BLOCK_SIZE_K // 2) + tl.arange(0, BLOCK_SIZE_K // 2)
    tl.store(fp4_ptr + offs_m[:, None] * stride_fm + offs_fk[None, :] * stride_fk, fp4_out)

    # scale: (BLOCK_SIZE_M, BLOCK_SIZE_K//QUANT_BLOCK_SIZE) — natural layout
    NUM_BLOCKS: tl.constexpr = BLOCK_SIZE_K // MXFP4_QUANT_BLOCK_SIZE
    offs_sk = pid_k * NUM_BLOCKS + tl.arange(0, NUM_BLOCKS)
    tl.store(scale_ptr + offs_m[:, None] * stride_sm + offs_sk[None, :] * stride_sk, scale_out)


# ── m=16 splitK: 只缓存 A 侧 buffer(shape 固定,省去反复 torch.empty)
# B 每次调用可能不同,view 是零拷贝,直接算即可,无需缓存
_m16_init = False
_m16_out  = None   # (M, N)    bf16  — pre-alloc output
_m16_A_q  = None   # (M, K//2) uint8 — pre-alloc quant A
_m16_A_sc = None   # (M, K//32) uint8 — pre-alloc quant A scale

# NUM_KSPLIT=7 → 17×7=119 CTAs
# get_splitk(K=7168, BSK=512, NKSPLIT=7) → SPLITK_BLOCK_SIZE=2048, actual NKSPLIT=7 ✓
_M16_CONFIG = {
    "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": 2, "matrix_instr_nonkdim": 16,
    "cache_modifier": ".cg", "NUM_KSPLIT": 7,
}
_M16_BSK = 512


def _init_m16_bufs(M, N, K, device):
    global _m16_init, _m16_out, _m16_A_q, _m16_A_sc
    _m16_out  = torch.empty((M, N),     dtype=dtypes.bf16,  device=device)
    _m16_A_q  = torch.empty((M, K // 2), dtype=torch.uint8, device=device)
    _m16_A_sc = torch.empty((M, K // 32), dtype=torch.uint8, device=device)
    _m16_init = True


def _run_m16_splitk(A, B_shuffle, B_scale_sh):
    global _m16_init
    M, K = A.shape
    N = B_shuffle.shape[0]

    if not _m16_init:
        _init_m16_bufs(M, N, K, A.device)

    # B 侧: 零拷贝 view,每次直接算(无计算开销)
    # B_scale_sh shape (sm, K//32) as fp8_e8m0,view 成 (sm//32, K) 与
    # _shuffle_scales 输出等价(K%256==0 时数学等价,已验证)
    sm = B_scale_sh.view(torch.uint8).shape[0]
    # 全部保持 uint8,不做 fp4x2/fp8_e8m0 view
    # benchmark 环境 Triton 不认识 float4_e2m1fn_x2 指针类型
    w      = B_shuffle.view(torch.uint8).reshape(N // 16, K // 2 * 16)
    wscale = B_scale_sh.view(torch.uint8).view(sm // 32, K)

    # Quant A → natural (M, K//32) scale, grid=(1,14)
    _mxfp4_quant_natural_kernel[
        (triton.cdiv(M, _M16_CONFIG["BLOCK_SIZE_M"]), triton.cdiv(K, _M16_BSK))
    ](
        A, _m16_A_q, _m16_A_sc,
        A.stride(0), A.stride(1),
        _m16_A_q.stride(0), _m16_A_q.stride(1),
        _m16_A_sc.stride(0), _m16_A_sc.stride(1),
        M=M, K=K,
        BLOCK_SIZE_M=_M16_CONFIG["BLOCK_SIZE_M"],
        BLOCK_SIZE_K=_M16_BSK,
        MXFP4_QUANT_BLOCK_SIZE=32,
        num_warps=4,
    )

    return gemm_afp4wfp4_preshuffle(
        _m16_A_q,
        w,
        _m16_A_sc,
        wscale,
        dtype=dtypes.bf16,
        y=_m16_out,
        config=dict(_M16_CONFIG),
        use_aot=False,
    )


# ── ASM path: small-M (m=4/32, k=512) ──────────────────────────────────────
_small_cache = {}

def _get_small_cache(M, K, N, device):
    key = (M, K, N)
    if key not in _small_cache:
        MXFP4_QUANT_BLOCK_SIZE = 32
        x_fp4      = torch.empty((M, K // 2), dtype=torch.uint8, device=device)
        scaleN     = K // MXFP4_QUANT_BLOCK_SIZE
        scaleN_pad = (scaleN + 7) // 8 * 8
        sm         = (M + 255) // 256 * 256
        bs_e8m0    = torch.empty(sm * scaleN_pad, dtype=torch.uint8, device=device)
        padded_m   = (M + 31) // 32 * 32
        gemm_out   = torch.empty((padded_m, N), dtype=dtypes.bf16, device=device)

        shape_cfg = SHAPE_CONFIGS.get((M, N, K), None)
        if shape_cfg is not None:
            kernel_name = _make_kernel_name(shape_cfg[0])
            log2_k_split = shape_cfg[1]
        else:
            kernel_name, log2_k_split = "", 0

        # k<=512: single-pass quant, small tile
        NUM_ITER, NUM_STAGES, NUM_WARPS = 1, 1, 4
        BLOCK_SIZE_N = max(32, min(256, triton.next_power_of_2(K)))
        BLOCK_SIZE_M = min(8, triton.next_power_of_2(M))
        grid = (triton.cdiv(M, BLOCK_SIZE_M), triton.cdiv(K, BLOCK_SIZE_N))

        _small_cache[key] = (
            x_fp4, bs_e8m0, gemm_out,
            scaleN, scaleN_pad, sm,
            kernel_name, log2_k_split,
            grid, BLOCK_SIZE_M, BLOCK_SIZE_N, NUM_ITER, NUM_STAGES, NUM_WARPS,
        )
    return _small_cache[key]


def _run_small(A, B_shuffle, B_scale_sh):
    M, K = A.shape
    N = B_shuffle.shape[0]
    (x_fp4, bs_e8m0, gemm_out,
     scaleN, scaleN_pad, sm,
     kernel_name, log2_k_split,
     grid, BLOCK_SIZE_M, BLOCK_SIZE_N, NUM_ITER, NUM_STAGES, NUM_WARPS,
    ) = _get_small_cache(M, K, N, A.device)

    _fused_mxfp4_quant_shuffle_kernel[grid](
        A, x_fp4, bs_e8m0,
        *A.stride(), *x_fp4.stride(),
        M=M, N=K, scaleN=scaleN, scaleN_pad=scaleN_pad,
        MXFP4_QUANT_BLOCK_SIZE=32, SCALING_MODE=0,
        NUM_ITER=NUM_ITER, BLOCK_SIZE_M=BLOCK_SIZE_M, BLOCK_SIZE_N=BLOCK_SIZE_N,
        NUM_STAGES=NUM_STAGES, num_warps=NUM_WARPS, waves_per_eu=0, num_stages=1,
    )
    A_q        = x_fp4.view(dtypes.fp4x2)
    A_scale_sh = bs_e8m0.view(sm, scaleN_pad).view(dtypes.fp8_e8m0)
    gemm_a4w4_asm(
        A_q, B_shuffle, A_scale_sh, B_scale_sh, gemm_out,
        kernel_name, None, 1.0, 0.0, True, log2_k_split,
    )
    return gemm_out[:M]


# ── Large-M path (m=64/256): CKGEMM ────────────────────────────────────────
def _quant_mxfp4_fused_simple(x):
    M, N = x.shape
    MXFP4_QUANT_BLOCK_SIZE = 32
    x_fp4 = torch.empty((M, N // 2), dtype=torch.uint8, device=x.device)
    scaleN = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
    scaleN_pad = (scaleN + 7) // 8 * 8
    sm = (M + 255) // 256 * 256
    bs_e8m0 = torch.empty(sm * scaleN_pad, dtype=torch.uint8, device=x.device)

    NUM_ITER, BLOCK_SIZE_M, BLOCK_SIZE_N, NUM_WARPS, NUM_STAGES = 4, 8, 128, 4, 2

    grid = (triton.cdiv(M, BLOCK_SIZE_M), triton.cdiv(N, BLOCK_SIZE_N * NUM_ITER))
    _fused_mxfp4_quant_shuffle_kernel[grid](
        x, x_fp4, bs_e8m0,
        *x.stride(), *x_fp4.stride(),
        M=M, N=N, scaleN=scaleN, scaleN_pad=scaleN_pad,
        MXFP4_QUANT_BLOCK_SIZE=MXFP4_QUANT_BLOCK_SIZE, SCALING_MODE=0,
        NUM_ITER=NUM_ITER, BLOCK_SIZE_M=BLOCK_SIZE_M, BLOCK_SIZE_N=BLOCK_SIZE_N,
        NUM_STAGES=NUM_STAGES, num_warps=NUM_WARPS, waves_per_eu=0, num_stages=1,
    )
    return x_fp4.view(dtypes.fp4x2), bs_e8m0.view(sm, scaleN_pad).view(dtypes.fp8_e8m0)


def _run_large(A, B_shuffle, B_scale_sh):
    A_q, A_scale_sh = _quant_mxfp4_fused_simple(A)
    return aiter.gemm_a4w4(
        A_q, B_shuffle, A_scale_sh, B_scale_sh,
        dtype=dtypes.bf16, bpreshuffle=True,
    )


# ── Dispatch ────────────────────────────────────────────────────────────────
def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    M, K = A.shape
    N = B_shuffle.shape[0]

    if M == 16 and K == 7168 and N == 2112:
        return _run_m16_splitk(A, B_shuffle, B_scale_sh)
    elif M <= 32:
        return _run_small(A, B_shuffle, B_scale_sh)
    else:
        return _run_large(A, B_shuffle, B_scale_sh)
scrolls · 318 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