Skip to content
KernelIndex
Search⌘K

submission 741444

Maxwell Cipher · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:aa2228fdabb9e2b0d3de70f7c598f2d3a2f5475cb335669b4c4840c830ede044
license declaredunknown
license concludedunknown
authorsMaxwell Cipher
imported2026-08-15

Techniques

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

num-warps = 4num_warps=4, waves_per_eu=0, num_stages=1,
split-kfrom aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
stages = 1NUM_ITER=1, NUM_STAGES=1,
tile-m = 16M=M, N=K, BLOCK_SIZE_M=16, BLOCK_SIZE_N=64,
tile-n = 64M=M, N=K, BLOCK_SIZE_M=16, BLOCK_SIZE_N=64,

Kernel source

MM_v33c.py187 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
from task import input_t, output_t
import torch
import triton
import triton.language as tl
from aiter import dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (
    _gemm_a16wfp4_preshuffle_kernel,
)
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
from aiter.ops.triton.gluon.gemm_afp4wfp4 import (
    _gemm_afp4wfp4_reduce_kernel as _reduce_partial,
)

_ASM = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"

# Shape 2: BK=512 (proven from v28)
_S2_SPLIT, _S2_BK, _S2_KS = get_splitk(3584, 512, 7)
_S2_GRID = (_S2_KS * triton.cdiv(16, 8) * triton.cdiv(2112, 128),)
_S2_AKS = triton.cdiv(3584, _S2_SPLIT // 2)
_S2_MKS = triton.next_power_of_2(_S2_KS)
_S2_REDUCE_GRID = (1, triton.cdiv(2112, 64))

# Precomputed configs for shapes 1/3/4/5: (bm, bn, bk, gm, nw, ns, wpe, cm, grid)
_DIRECT_CFG: dict[tuple[int, int, int], tuple] = {
    (4, 2880, 512):    (4, 128, 256, 1, 4, 2, 0, ".cg", 23),
    (32, 4096, 512):   (8, 128, 256, 1, 4, 2, 2, None, 128),
    (32, 2880, 512):   (8, 128, 256, 1, 4, 2, 2, None, 92),
    (64, 7168, 2048):  (16, 128, 256, 1, 4, 2, 2, ".cg", 224),
}


@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 _quant_shuffle_kernel(
    x_ptr, fp4_ptr, sc_ptr,
    stride_x_m_in, stride_x_n_in, stride_f_m_in, stride_f_n_in,
    M, N,
    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, SCALE_N_PAD: tl.constexpr,
    SCALE_STORE_WT: 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_f_m = tl.cast(stride_f_m_in, tl.int64)
    stride_f_n = tl.cast(stride_f_n_in, tl.int64)
    nq: 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):
        offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
        x_idx = offs_m[:, None] * stride_x_m + offs_n[None, :] * stride_x_n
        if EVEN_M_N:
            x = tl.load(x_ptr + x_idx, cache_modifier=".cg").to(tl.float32)
        else:
            mask = (offs_m < M)[:, None] & (offs_n < N)[None, :]
            x = tl.load(x_ptr + x_idx, mask=mask, cache_modifier=".cg").to(tl.float32)
        fp4_out, sc_out = _mxfp4_quant_op(x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE)
        fp4_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        fp4_n = pid_n * (BLOCK_SIZE_N // 2) + tl.arange(0, BLOCK_SIZE_N // 2)
        fp4_idx = fp4_m[:, None] * stride_f_m + fp4_n[None, :] * stride_f_n
        if EVEN_M_N:
            tl.store(fp4_ptr + fp4_idx, fp4_out, cache_modifier=".wt")
        else:
            fp4_mask = (fp4_m < M)[:, None] & (fp4_n < (N // 2))[None, :]
            tl.store(fp4_ptr + fp4_idx, fp4_out, mask=fp4_mask, cache_modifier=".wt")
        sc_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
        sc_n = pid_n * nq + tl.arange(0, nq)
        total_sc_cols = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
        t0 = sc_m[:, None] // 32
        t1 = sc_m[:, None] % 32
        t2 = t1 % 16
        t1 = t1 // 16
        t3 = sc_n[None, :] // 8
        t4 = sc_n[None, :] % 8
        t5 = t4 % 4
        t4 = t4 // 4
        sc_idx = (t1 + t4 * 2 + t2 * 4 + t5 * 64 + t3 * 256 + t0 * 32 * SCALE_N_PAD)
        valid_sc = (sc_m < M)[:, None] & (sc_n < total_sc_cols)[None, :]
        sc_out = tl.where(valid_sc, sc_out, 127)
        scale_rows = (M + 255) // 256 * 256
        scale_mask = (sc_m < scale_rows)[:, None] & (sc_n < SCALE_N_PAD)[None, :]
        if SCALE_STORE_WT:
            tl.store(sc_ptr + sc_idx, sc_out.to(tl.uint8), mask=scale_mask, cache_modifier=".wt")
        else:
            tl.store(sc_ptr + sc_idx, sc_out.to(tl.uint8), mask=scale_mask, cache_modifier=".cg")


def custom_kernel(data: input_t) -> output_t:
    A = data[0]
    B_sh = data[3]
    B_sc = data[4]
    M, K = A.shape
    N = B_sh.shape[0]

    # Shape 6: baseline
    if M > 64:
        x_fp4 = torch.empty((M, K >> 1), dtype=torch.uint8, device=A.device)
        x_sc = torch.empty((256, 48), dtype=torch.uint8, device=A.device)
        out = torch.empty((256, N), dtype=torch.bfloat16, device=A.device)
        _quant_shuffle_kernel[(16, 24)](
            A, x_fp4, x_sc,
            A.stride(0), A.stride(1), x_fp4.stride(0), x_fp4.stride(1),
            M=M, N=K, BLOCK_SIZE_M=16, BLOCK_SIZE_N=64,
            NUM_ITER=1, NUM_STAGES=1,
            MXFP4_QUANT_BLOCK_SIZE=32, SCALING_MODE=0,
            SCALE_N_PAD=48, SCALE_STORE_WT=True,
            num_warps=4, waves_per_eu=0, num_stages=1,
        )
        gemm_a4w4_asm(
            x_fp4.view(dtypes.fp4x2), B_sh,
            x_sc.view(dtypes.fp8_e8m0), B_sc,
            out, _ASM, None, 1.0, 0.0, True, log2_k_split=1,
        )
        return out[:M]

    # Reshape B for preshuffle kernel
    pw = B_sh.view(torch.uint8).reshape(N >> 4, (K >> 1) << 4)
    s0, s1 = B_sc.shape
    ps = B_sc.view(torch.uint8).reshape(s0 >> 5, s1 << 5)
    pk = K >> 1

    # *** Shape 2: BK=512, warps=4 (proven from v28) ***
    if K > 4096:
        out = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
        part = torch.empty((_S2_KS, M, N), dtype=torch.float32, device=A.device)
        _gemm_a16wfp4_preshuffle_kernel[_S2_GRID](
            A, pw, part, ps, M, N, pk,
            A.stride(0), A.stride(1), pw.stride(0), pw.stride(1),
            part.stride(0), part.stride(1), part.stride(2),
            ps.stride(0), ps.stride(1),
            BLOCK_SIZE_M=8, BLOCK_SIZE_N=128, BLOCK_SIZE_K=_S2_BK,
            GROUP_SIZE_M=1, NUM_KSPLIT=_S2_KS,
            SPLITK_BLOCK_SIZE=_S2_SPLIT,
            num_warps=4, num_stages=2, waves_per_eu=2,
            matrix_instr_nonkdim=16, PREQUANT=True, cache_modifier=".cg",
        )
        _reduce_partial[_S2_REDUCE_GRID](
            part, out, M, N,
            part.stride(0), part.stride(1), part.stride(2),
            out.stride(0), out.stride(1),
            16, 64, _S2_AKS, _S2_MKS,
        )
        return out

    # Shapes 1/3/4/5: precomputed configs with fallback
    cfg = _DIRECT_CFG.get((M, N, K))
    if cfg is not None:
        bm, bn, bk, gm, nw, ns, wpe, cm, gs = cfg
    else:
        if M <= 4:
            bm, bn, bk, gm, nw, ns, wpe, cm = 4, 128, 256, 1, 4, 2, 0, ".cg"
        elif M <= 8:
            bm, bn, bk, gm, nw, ns, wpe, cm = 8, 128, 256, 1, 4, 2, 0, ".cg"
        elif M <= 32 and K <= 1024:
            bm, bn, bk, gm, nw, ns, wpe, cm = 8, 128, 256, 1, 4, 2, 2, None
        elif M <= 32:
            bm, bn, bk, gm, nw, ns, wpe, cm = 32, 64, 512, 1, 8, 1, 2, None
        else:
            bm, bn, bk, gm, nw, ns, wpe, cm = 16, 128, 256, 1, 4, 2, 2, ".cg"
        bn = max(bn, 32)
        gs = triton.cdiv(M, bm) * triton.cdiv(N, bn)
    out = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)

    _gemm_a16wfp4_preshuffle_kernel[(gs,)](
        A, pw, out, ps, M, N, pk,
        A.stride(0), A.stride(1), pw.stride(0), pw.stride(1),
        0, out.stride(0), out.stride(1),
        ps.stride(0), ps.stride(1),
        BLOCK_SIZE_M=bm, BLOCK_SIZE_N=bn, BLOCK_SIZE_K=bk,
        GROUP_SIZE_M=gm, NUM_KSPLIT=1, SPLITK_BLOCK_SIZE=2 * pk,
        num_warps=nw, num_stages=ns, waves_per_eu=wpe,
        matrix_instr_nonkdim=16, PREQUANT=True, cache_modifier=cm,
    )
    return out
scrolls · 187 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