Skip to content
KernelIndex
Search⌘K

submission 736751

sepehresy · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:ec79f54f24afb15a856f6f80d34f85b08d2ccd20b8b8931f526e2d02242b1d84
license declaredunknown
license concludedunknown
authorssepehresy
imported2026-08-26

Techniques

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

tile-k = 512BLOCK_K = 512
tile-m = 16BLOCK_M = 16
tile-n = 128BLOCK_N = 128

Kernel source

MM_V68.py238 lines
"""
MM_V68: V67 + deeper pipeline for the main K=1536 M=256 bottleneck.

V67 fixed KernelGuard issues and improved the K=512/K=7168 rows, but the
ranked leaderboard still shows the slow tail centered on:
  (256, 3072, 1536) ~17.8us

This version keeps V67's KernelGuard-safe structure and only bumps
num_stages from 2 -> 3 for that one shape to test whether the 3 x BLOCK_K
pipeline benefits from deeper software pipelining without changing the
winning configs elsewhere.
"""
import os
os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
os.environ["GPU_MAX_HW_QUEUES"] = "2"

import torch
import triton
import triton.language as tl
from task import input_t, output_t

from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import _mxfp4_quant_op

_u8 = torch.uint8
_bf16 = torch.bfloat16
_f32 = torch.float32


@triton.jit
def _fused_gemm_kernel(
    a_ptr, b_ptr, c_ptr, bs_ptr,
    M, N, K,
    stride_am, stride_ak,
    stride_bk, stride_bn,
    stride_ck, stride_cm, stride_cn,
    SN,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
    GROUP_SIZE_M: tl.constexpr, NUM_KSPLIT: tl.constexpr,
    K_ITERS_PER_SPLIT: tl.constexpr,
    EVEN_MNK: tl.constexpr,
    num_warps: tl.constexpr, num_stages: tl.constexpr,
    waves_per_eu: tl.constexpr,
):
    K_HALF: tl.constexpr = BLOCK_K // 2
    SCALE_K: tl.constexpr = BLOCK_K // 32
    SN32 = SN * 32

    GRID_MN = tl.cdiv(M, BLOCK_M) * tl.cdiv(N, BLOCK_N)
    pid_unified = tl.program_id(0)
    pid_k = pid_unified // GRID_MN
    pid = pid_unified % GRID_MN

    num_pid_m = tl.cdiv(M, BLOCK_M)
    num_pid_n = tl.cdiv(N, BLOCK_N)

    if NUM_KSPLIT == 1:
        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 = tl.minimum(num_pid_m - first_pid_m, GROUP_SIZE_M)
        pid_m = first_pid_m + ((pid % num_pid_in_group) % group_size_m)
        pid_n = (pid % num_pid_in_group) // group_size_m
    else:
        pid_m = pid // num_pid_n
        pid_n = pid % num_pid_n

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)

    sn_i0 = offs_n // 32
    sn_i1 = (offs_n >> 4) & 1
    sn_i2 = offs_n & 15
    src_r = sn_i0 * SN32 + sn_i1 + sn_i2 * 4

    a_base = a_ptr + offs_m[:, None] * stride_am

    accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
    k_start_split = pid_k * K_ITERS_PER_SPLIT * BLOCK_K

    for ki in range(K_ITERS_PER_SPLIT):
        k_start = k_start_split + ki * BLOCK_K

        if k_start < K:
            k_half_start = k_start // 2
            a_k_offs = k_start + tl.arange(0, BLOCK_K)
            b_k_offs = k_half_start + tl.arange(0, BLOCK_K // 2)

            if EVEN_MNK:
                a_tile = tl.load(a_base + a_k_offs[None, :] * stride_ak)
                b_tile = tl.load(b_ptr + b_k_offs[:, None] * stride_bk + offs_n[None, :] * stride_bn)
            else:
                a_mask = (offs_m[:, None] < M) & (a_k_offs[None, :] < K)
                a_tile = tl.load(a_base + a_k_offs[None, :] * stride_ak, mask=a_mask, other=0.0)
                b_mask = (b_k_offs[:, None] < (K // 2)) & (offs_n[None, :] < N)
                b_tile = tl.load(
                    b_ptr + b_k_offs[:, None] * stride_bk + offs_n[None, :] * stride_bn,
                    mask=b_mask,
                    other=0,
                )

            scale_k_idx = k_start // 32
            sk_offs = scale_k_idx + tl.arange(0, SCALE_K)
            sk_i3 = sk_offs // 8
            sk_i4 = (sk_offs >> 2) & 1
            sk_i5 = sk_offs & 3
            src_c = sk_i3 * 256 + sk_i4 * 2 + sk_i5 * 64
            bs_src = src_r[:, None] + src_c[None, :]

            if EVEN_MNK:
                b_scales = tl.load(bs_ptr + bs_src)
            else:
                bs_mask = (offs_n[:, None] < N) & (sk_offs[None, :] < (K // 32))
                b_scales = tl.load(bs_ptr + bs_src, mask=bs_mask, other=0)

            a_fp4, a_scales = _mxfp4_quant_op(a_tile, BLOCK_K, BLOCK_M, 32)
            accumulator = tl.dot_scaled(a_fp4, a_scales, "e2m1", b_tile, b_scales, "e2m1", accumulator)

    c = accumulator.to(tl.bfloat16) if NUM_KSPLIT == 1 else accumulator
    c_mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
    c_ptrs = c_ptr + pid_k * stride_ck + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
    tl.store(c_ptrs, c, mask=c_mask)


@triton.jit
def _reduce_kernel(
    partials_ptr, out_ptr, M, N,
    stride_pk, stride_pm, stride_pn,
    stride_om, stride_on,
    NUM_SPLITS: 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 k in range(NUM_SPLITS):
        p = tl.load(
            partials_ptr + k * stride_pk + offs_m[:, None] * stride_pm + offs_n[None, :] * stride_pn,
            mask=mask,
            other=0.0,
        )
        acc += p
    tl.store(
        out_ptr + offs_m[:, None] * stride_om + offs_n[None, :] * stride_on,
        acc.to(tl.bfloat16),
        mask=mask,
    )


# (BLOCK_M, NUM_KSPLIT, K_ITERS, num_warps, num_stages, EVEN_MNK)
_SHAPE_CFGS = {
    (4, 2880, 512):    (16, 1, 1, 8, 2, False),
    (32, 4096, 512):   (16, 1, 1, 8, 2, True),
    (32, 2880, 512):   (16, 1, 1, 8, 2, False),
    (256, 2880, 512):  (16, 1, 1, 8, 2, False),
    (8, 2112, 7168):   (16, 14, 1, 4, 1, False),
    (16, 2112, 7168):  (16, 14, 1, 4, 1, False),
    (64, 7168, 2048):  (16, 2, 2, 4, 2, True),
    (256, 3072, 1536): (16, 1, 3, 4, 3, True),
    (16, 3072, 1536):  (16, 1, 3, 4, 2, True),
    (64, 3072, 1536):  (16, 1, 3, 4, 2, True),
}


def custom_kernel(data: input_t) -> output_t:
    A = data[0]
    B_q = data[2]
    B_scale_sh = data[4]

    if not A.is_contiguous():
        A = A.contiguous()

    m, k = A.shape
    n = B_q.shape[0]

    B_u8 = B_q.view(_u8)
    bs_u8 = B_scale_sh.view(_u8)
    sn = bs_u8.shape[1]

    BLOCK_N = 128
    BLOCK_K = 512

    cfg = _SHAPE_CFGS.get((m, n, k))
    if cfg is not None:
        BLOCK_M, NUM_KSPLIT, K_ITERS, nw, ns, even = cfg
    else:
        BLOCK_M = 16
        NUM_KSPLIT = max(1, k // BLOCK_K) if k >= 4096 else 1
        K_ITERS = max(1, triton.cdiv(k, max(NUM_KSPLIT, 1) * BLOCK_K))
        nw, ns, even = 4, 1, False

    grid_mn = triton.cdiv(m, BLOCK_M) * triton.cdiv(n, BLOCK_N)

    if NUM_KSPLIT <= 1:
        out = torch.empty((m, n), dtype=_bf16, device=A.device)
        _fused_gemm_kernel[(grid_mn,)](
            A, B_u8, out, bs_u8,
            m, n, k,
            A.stride(0), A.stride(1),
            B_u8.stride(1), B_u8.stride(0),
            0, out.stride(0), out.stride(1),
            sn,
            BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
            GROUP_SIZE_M=4, NUM_KSPLIT=1,
            K_ITERS_PER_SPLIT=K_ITERS,
            EVEN_MNK=even,
            num_warps=nw, num_stages=ns, waves_per_eu=2,
        )
        return out

    partials = torch.empty((NUM_KSPLIT, m, n), dtype=_f32, device=A.device)
    out = torch.empty((m, n), dtype=_bf16, device=A.device)

    _fused_gemm_kernel[(NUM_KSPLIT * grid_mn,)](
        A, B_u8, partials, bs_u8,
        m, n, k,
        A.stride(0), A.stride(1),
        B_u8.stride(1), B_u8.stride(0),
        partials.stride(0), partials.stride(1), partials.stride(2),
        sn,
        BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
        GROUP_SIZE_M=4, NUM_KSPLIT=NUM_KSPLIT,
        K_ITERS_PER_SPLIT=K_ITERS,
        EVEN_MNK=even,
        num_warps=nw, num_stages=ns, waves_per_eu=2,
    )

    _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_SPLITS=NUM_KSPLIT,
        BLOCK_M=16, BLOCK_N=64,
    )
    return out
scrolls · 238 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