Skip to content
KernelIndex
Search⌘K

submission 708351

ihansel · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:beeb3a2c0f9736b03fb7a60f33d5059c6b877f4817b3ca467d6493171366b514
license declaredunknown
license concludedunknown
authorsihansel
imported2026-08-26

Techniques

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

fp4"""MXFP4 GEMM v129: Hybrid best-of-both.
num-warps = 4NUM_KSPLIT=NUM_KSPLIT, BLOCK_M=16, BLOCK_N=64, num_warps=4,
split-ksplitK = 0
stages = 2num_warps=num_warps, num_stages=2, matrix_instr_nonkdim=16,
tile-m = 16NUM_KSPLIT=NUM_KSPLIT, BLOCK_M=16, BLOCK_N=64, num_warps=4,
tile-n = 64NUM_KSPLIT=NUM_KSPLIT, BLOCK_M=16, BLOCK_N=64, num_warps=4,

Kernel source

submission.py338 lines
"""MXFP4 GEMM v129: Hybrid best-of-both.
M<=32: AITER quant fused into Triton GEMM (v128, 7.3-10.4µs)
M>=64: fused quant+shuffle + direct ASM (v126, 15.9-17.2µs)
"""
import torch
import triton
import triton.language as tl
import aiter  # noqa: F401
from aiter import dtypes
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm, get_GEMM_config

from task import input_t, output_t
from reference import ref_kernel  # noqa: F401


# ==================== FUSED QUANT+SHUFFLE KERNEL (for M>=64 ASM path) ====================
@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_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: tl.constexpr, N: tl.constexpr, scaleN_valid: tl.constexpr, scaleN_pad: tl.constexpr,
    MXFP4_QUANT_BLOCK_SIZE: tl.constexpr, SCALING_MODE: tl.constexpr,
    NUM_ITER: tl.constexpr, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
    NUM_STAGES: tl.constexpr, EVEN_M_N: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)
    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 iter in tl.static_range(NUM_ITER):
        offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        offs_n = (pid_n * NUM_ITER + iter) * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
        x_offs = offs_m[:, None] * stride_x_m + offs_n[None, :] * stride_x_n
        if EVEN_M_N:
            x = tl.load(x_ptr + x_offs).to(tl.float32)
        else:
            x_mask = (offs_m < M)[:, None] & (offs_n < N)[None, :]
            x = tl.load(x_ptr + x_offs, mask=x_mask, other=0.0).to(tl.float32)
        x_fp4, bs_e8m0 = _mxfp4_quant_op(x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE)
        offs_fp4_n = (pid_n * NUM_ITER + iter) * (BLOCK_SIZE_N // 2) + tl.arange(0, BLOCK_SIZE_N // 2)
        out_offs = offs_m[:, None] * stride_x_fp4_m + offs_fp4_n[None, :] * stride_x_fp4_n
        if EVEN_M_N:
            tl.store(x_fp4_ptr + out_offs, x_fp4)
        else:
            out_mask = (offs_m < M)[:, None] & (offs_fp4_n < (N // 2))[None, :]
            tl.store(x_fp4_ptr + out_offs, x_fp4, mask=out_mask)
        bs_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        bs_col_base = (pid_n * NUM_ITER + iter) * NUM_QUANT_BLOCKS
        bs_offs_n = bs_col_base + tl.arange(0, NUM_QUANT_BLOCKS)
        g0 = bs_offs_m // 32
        rem32 = bs_offs_m % 32
        g1 = rem32 // 16
        g2 = rem32 % 16
        g3 = bs_offs_n // 8
        g4 = (bs_offs_n % 8) // 4
        g5 = bs_offs_n % 4
        shuffled_offs = (g1[:, None] + g4[None, :] * 2 + g2[:, None] * 4
                         + g5[None, :] * 64 + g3[None, :] * 256
                         + g0[:, None] * (scaleN_pad * 32))
        bs_mask_valid = (bs_offs_m < M)[:, None] & (bs_offs_n < scaleN_valid)[None, :]
        bs_e8m0_padded = tl.where(bs_mask_valid, bs_e8m0, 127)
        sm_pad = ((M + 255) // 256) * 256
        bs_mask_write = (bs_offs_m < sm_pad)[:, None] & (bs_offs_n < scaleN_pad)[None, :]
        tl.store(bs_ptr + shuffled_offs, bs_e8m0_padded, mask=bs_mask_write)


@triton.jit
def _aiter_fused_quant_gemm_kernel(
    a_bf16_ptr, b_ptr, c_ptr, b_scales_ptr,
    M, N, K,
    stride_am, stride_ak, stride_bk, stride_bn,
    stride_cs, stride_cm, stride_cn,
    stride_bsn, stride_bsk,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
    NUM_KSPLIT: tl.constexpr,
    MXFP4_QUANT_BLOCK_SIZE: tl.constexpr,
):
    SCALE_GROUP_SIZE: tl.constexpr = MXFP4_QUANT_BLOCK_SIZE

    pid = tl.program_id(axis=0)
    num_pid_m = tl.cdiv(M, BLOCK_M)
    num_pid_n = tl.cdiv(N, BLOCK_N)
    num_pid_mn = num_pid_m * num_pid_n
    pid_k = pid // num_pid_mn
    pid_mn = pid % num_pid_mn
    pid_m = pid_mn // num_pid_n
    pid_n = pid_mn % num_pid_n

    total_k_iters = tl.cdiv(K, BLOCK_K)
    k_iters_per_split = tl.cdiv(total_k_iters, NUM_KSPLIT)
    k_start = pid_k * k_iters_per_split
    k_end = min((pid_k + 1) * k_iters_per_split, total_k_iters)

    offs_am = (pid_m * BLOCK_M + tl.arange(0, BLOCK_M)) % M
    offs_k_bf16 = tl.arange(0, BLOCK_K)
    a_bf16_ptrs = a_bf16_ptr + (offs_am[:, None] * stride_am + offs_k_bf16[None, :] * stride_ak)
    a_bf16_ptrs += k_start * BLOCK_K * stride_ak

    offs_k_fp4 = tl.arange(0, BLOCK_K // 2)
    offs_bn = (pid_n * BLOCK_N + tl.arange(0, BLOCK_N)) % N
    b_ptrs = b_ptr + (offs_k_fp4[:, None] * stride_bk + offs_bn[None, :] * stride_bn)
    b_ptrs += k_start * (BLOCK_K // 2) * stride_bk

    offs_ks = tl.arange(0, BLOCK_K // SCALE_GROUP_SIZE * 32)
    offs_bsn = (pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)) % tl.cdiv(N, 32)
    b_scale_ptrs = b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk
    b_scale_ptrs += k_start * BLOCK_K * stride_bsk

    accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    for k in range(k_start, k_end):
        # Load bf16 A tile
        a_bf16 = tl.load(a_bf16_ptrs)
        a_f32 = a_bf16.to(tl.float32)

        # Use AITER's official quant function — exact same as reference
        a_fp4, a_scales = _mxfp4_quant_op(a_f32, BLOCK_K, BLOCK_M, MXFP4_QUANT_BLOCK_SIZE)

        # Load B (fp4, pre-quantized)
        b = tl.load(b_ptrs, cache_modifier=".cg")
        b_scales = tl.load(b_scale_ptrs).reshape(
            BLOCK_N // 32, BLOCK_K // SCALE_GROUP_SIZE // 8, 4, 16, 2, 2, 1
        ).permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_N, BLOCK_K // SCALE_GROUP_SIZE)

        # MXFP4 × MXFP4 GEMM via tl.dot_scaled
        accumulator = tl.dot_scaled(a_fp4, a_scales, "e2m1", b, b_scales, "e2m1",
                                     acc=accumulator, fast_math=True)

        a_bf16_ptrs += BLOCK_K * stride_ak
        b_ptrs += (BLOCK_K // 2) * stride_bk
        b_scale_ptrs += BLOCK_K * stride_bsk

    offs_cm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M).to(tl.int64)
    offs_cn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N).to(tl.int64)
    c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
    if NUM_KSPLIT == 1:
        c = accumulator.to(c_ptr.type.element_ty)
        c_ptrs = c_ptr + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :]
        tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")
    else:
        c_ptrs = c_ptr + pid_k * stride_cs + stride_cm * offs_cm[:, None] + stride_cn * offs_cn[None, :]
        tl.store(c_ptrs, accumulator, mask=c_mask)


@triton.jit
def _reduce_kernel(
    partials_ptr, out_ptr, M, N,
    stride_ps, stride_pm, stride_pn, stride_om, stride_on,
    NUM_KSPLIT: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
):
    pid_m = tl.program_id(0)
    pid_n = tl.program_id(1)
    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
    for s in range(NUM_KSPLIT):
        ptrs = partials_ptr + s * stride_ps + offs_m[:, None] * stride_pm + offs_n[None, :] * stride_pn
        acc += tl.load(ptrs, mask=mask, other=0.0)
    out_ptrs = out_ptr + offs_m[:, None] * stride_om + offs_n[None, :] * stride_on
    tl.store(out_ptrs, acc.to(out_ptr.type.element_ty), mask=mask)


_buffers = {}


def _get_config(m, n, k):
    """Shape-specialized configs. Fused Triton for ALL shapes."""
    if m <= 16:
        if k >= 4096:
            return 16, 128, 256, 14, 8
        elif k >= 1536:
            return 16, 128, 256, 4, 8
        else:
            return 16, 128, 256, 1, 8
    elif m <= 32:
        if k >= 4096:
            return 16, 64, 256, 8, 8
        elif k >= 1536:
            return 16, 64, 256, 4, 8
        else:
            return 16, 64, 256, 1, 8
    elif m <= 64:
        # M=64: BLOCK_M=32, 2 M-tiles. Fused quant eliminates 10µs overhead.
        if k >= 4096:
            return 32, 128, 256, 4, 8
        elif k >= 1536:
            return 32, 128, 256, 2, 8
        else:
            return 32, 128, 256, 1, 8
    else:
        # M=256: BLOCK_M=32, 8 M-tiles. Large enough for good CU fill.
        if k >= 4096:
            return 32, 128, 256, 2, 8
        elif k >= 1536:
            return 32, 128, 256, 1, 8
        else:
            return 32, 128, 256, 1, 8


def _triton_dispatch(A_bf16, B_q, B_scale_sh, m, n, k):
    BLOCK_M, BLOCK_N, BLOCK_K, NUM_KSPLIT, num_warps = _get_config(m, n, k)
    B_q_u8 = B_q.view(torch.uint8)
    B_t = B_q_u8.T
    B_s = B_scale_sh.view(torch.uint8)
    B_scales_triton = B_s.reshape(B_s.shape[0] // 32, B_s.shape[1] * 32)
    num_pid_m = triton.cdiv(m, BLOCK_M)
    num_pid_n = triton.cdiv(n, BLOCK_N)

    key = (m, n, k)
    if key not in _buffers:
        if NUM_KSPLIT == 1:
            _buffers[key] = (
                torch.empty((m, n), dtype=torch.bfloat16, device=A_bf16.device),
                None,
            )
        else:
            _buffers[key] = (
                torch.empty((m, n), dtype=torch.bfloat16, device=A_bf16.device),
                torch.empty((NUM_KSPLIT, m, n), dtype=torch.float32, device=A_bf16.device),
            )
    out, partials = _buffers[key]

    if NUM_KSPLIT == 1:
        _aiter_fused_quant_gemm_kernel[(num_pid_m * num_pid_n,)](
            A_bf16, B_t, out, B_scales_triton, m, n, k,
            A_bf16.stride(0), A_bf16.stride(1),
            B_t.stride(0), B_t.stride(1),
            0, out.stride(0), out.stride(1),
            B_scales_triton.stride(0), B_scales_triton.stride(1),
            BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
            NUM_KSPLIT=1, MXFP4_QUANT_BLOCK_SIZE=32,
            num_warps=num_warps, num_stages=2, matrix_instr_nonkdim=16,
        )
    else:
        _aiter_fused_quant_gemm_kernel[(NUM_KSPLIT * num_pid_m * num_pid_n,)](
            A_bf16, B_t, partials, B_scales_triton, m, n, k,
            A_bf16.stride(0), A_bf16.stride(1),
            B_t.stride(0), B_t.stride(1),
            partials.stride(0), partials.stride(1), partials.stride(2),
            B_scales_triton.stride(0), B_scales_triton.stride(1),
            BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
            NUM_KSPLIT=NUM_KSPLIT, MXFP4_QUANT_BLOCK_SIZE=32,
            num_warps=num_warps, num_stages=2, matrix_instr_nonkdim=16,
        )
        _reduce_kernel[(triton.cdiv(m, 16), triton.cdiv(n, 64))](
            partials, out, m, n,
            partials.stride(0), partials.stride(1), partials.stride(2),
            out.stride(0), out.stride(1),
            NUM_KSPLIT=NUM_KSPLIT, BLOCK_M=16, BLOCK_N=64, num_warps=4,
        )
    return out


_asm_buffers = {}
MXFP4_QBS = 32


def _asm_path(A, B_shuffle, B_scale_sh):
    """v126's fast path: fused quant+shuffle + direct ASM GEMM."""
    M, K = A.shape
    N = B_shuffle.shape[0]
    key = (M, N, K)
    if key not in _asm_buffers:
        scaleN_valid = triton.cdiv(K, MXFP4_QBS)
        scaleN_pad = triton.cdiv(scaleN_valid, 8) * 8
        sm_pad = triton.cdiv(M, 256) * 256
        padded_m = (M + 31) // 32 * 32
        x_fp4 = torch.empty((M, K // 2), dtype=torch.uint8, device=A.device)
        scale_sh = torch.full((sm_pad * scaleN_pad,), 127, dtype=torch.uint8, device=A.device)
        out = torch.empty((padded_m, N), dtype=dtypes.bf16, device=A.device)
        ck_config = get_GEMM_config(M, N, K)
        splitK = 0
        kernelName = ""
        if ck_config is not None:
            splitK = ck_config.get("splitK", None)
            splitK = 0 if splitK is None else splitK
            kernelName = ck_config["kernelName"]
        NUM_ITER, BSM, BSN, NW, NS = _get_quant_config(M, K)
        _asm_buffers[key] = (x_fp4, scale_sh, scaleN_valid, scaleN_pad, sm_pad,
                             out, padded_m, kernelName, splitK, NUM_ITER, BSM, BSN, NW, NS)
    (x_fp4, scale_sh_flat, scaleN_valid, scaleN_pad, sm_pad,
     out, padded_m, kernelName, splitK, NUM_ITER, BSM, BSN, NW, NS) = _asm_buffers[key]

    grid = (triton.cdiv(M, BSM), triton.cdiv(K, BSN * NUM_ITER))
    _fused_quant_shuffle_kernel[grid](
        A, x_fp4, scale_sh_flat, *A.stride(), *x_fp4.stride(),
        M=M, N=K, scaleN_valid=scaleN_valid, scaleN_pad=scaleN_pad,
        MXFP4_QUANT_BLOCK_SIZE=MXFP4_QBS, SCALING_MODE=0,
        NUM_ITER=NUM_ITER, BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=BSN,
        NUM_STAGES=NS, num_warps=NW, waves_per_eu=0, num_stages=1,
    )
    A_q = x_fp4.view(dtypes.fp4x2)
    A_scale_sh = scale_sh_flat.view(sm_pad, scaleN_pad).view(dtypes.fp8_e8m0)
    gemm_a4w4_asm(A_q.view(M, K // 2), B_shuffle, A_scale_sh, B_scale_sh,
                   out, kernelName, None, 1.0, 0.0, True, log2_k_split=splitK)
    return out[:M]


def _get_quant_config(M, N):
    """Same config logic as aiter.ops.triton.quant.dynamic_mxfp4_quant."""
    if M <= 32:
        NUM_ITER = 1
        BLOCK_SIZE_M = triton.next_power_of_2(M)
        BLOCK_SIZE_N = 32
        NUM_WARPS = 1
        NUM_STAGES = 1
    else:
        NUM_ITER = 4
        BLOCK_SIZE_M = 64
        BLOCK_SIZE_N = 64
        NUM_WARPS = 4
        NUM_STAGES = 2
        if N <= 16384:
            BLOCK_SIZE_M = 32
            BLOCK_SIZE_N = 128
    if N <= 1024:
        NUM_ITER = 1
        NUM_STAGES = 1
        NUM_WARPS = 4
        BLOCK_SIZE_N = min(256, triton.next_power_of_2(N))
        BLOCK_SIZE_N = max(32, BLOCK_SIZE_N)
        BLOCK_SIZE_M = min(8, triton.next_power_of_2(M))
    return NUM_ITER, BLOCK_SIZE_M, BLOCK_SIZE_N, NUM_WARPS, NUM_STAGES


def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    m, k = A.shape
    n = B.shape[0]
    return _triton_dispatch(A, B_q, B_scale_sh, m, n, k)
scrolls · 338 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