Skip to content
KernelIndex
Search⌘K

submission 690748

sean_nobricks · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:8e62404b1108f3ce5c92367bb9ccde518db38fd6c3da7a5b235b61399b2ca9d6
license declaredunknown
license concludedunknown
authorssean_nobricks
imported2026-08-15

Techniques

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

fp4"""MXFP4 GEMM — custom Triton kernel with in-kernel scale unshuffle and split-K."""
num-warps = 4num_warps=4, num_stages=2,
split-k"""MXFP4 GEMM — custom Triton kernel with in-kernel scale unshuffle and split-K."""
stages = 2num_warps=4, num_stages=2,
tile-k = 32BLOCK_M, BLOCK_N, BLOCK_K = 32, 32, 256

Kernel source

submission.py150 lines
"""MXFP4 GEMM — custom Triton kernel with in-kernel scale unshuffle and split-K."""
import torch
import triton
import triton.language as tl
from task import input_t, output_t

SCALE_GROUP_SIZE = 32


@triton.jit
def mxfp4_gemm_splitk_kernel(
    a_ptr, b_ptr, c_ptr, a_scale_ptr, b_scale_ptr,
    M, N, K_packed,
    stride_am, stride_ak,
    stride_bk, stride_bn,
    stride_cm, stride_cn,
    stride_asm, stride_ask,
    stride_bsn, stride_bsk,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
    SPLIT_K: tl.constexpr,
):
    SG: tl.constexpr = 32

    pid_mn = tl.program_id(0)
    pid_k = tl.program_id(1)

    num_pid_n = tl.cdiv(N, BLOCK_N)
    pid_m = pid_mn // num_pid_n
    pid_n = pid_mn % num_pid_n

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

    k_per_split = tl.cdiv(K_packed, SPLIT_K * (BLOCK_K // 2)) * (BLOCK_K // 2)
    k_start = pid_k * k_per_split
    k_end = tl.minimum(k_start + k_per_split, K_packed)

    offs_k = tl.arange(0, BLOCK_K // 2)

    a_ptrs = a_ptr + offs_m[:, None] * stride_am + (k_start + offs_k[None, :]) * stride_ak
    b_ptrs = b_ptr + (k_start + offs_k[:, None]) * stride_bk + offs_n[None, :] * stride_bn

    # A scales: un-shuffled natural (M, K_scale) layout
    num_scale_k: tl.constexpr = BLOCK_K // SG
    offs_sk = tl.arange(0, num_scale_k)
    scale_k_start = k_start * 2 // SG
    a_scale_ptrs = a_scale_ptr + offs_m[:, None] * stride_asm + (scale_k_start + offs_sk[None, :]) * stride_ask

    # B scales: shuffled (N//32, K_scale*32) layout per Triton CDNA4 tutorial.
    # Load contiguous shuffled block, reshape/permute in-register to recover
    # logical (BLOCK_N, BLOCK_K//SG) layout. Compiler detects this pattern
    # and enables 4x vectorized scale loads.
    SHUFFLED_SCALE_K: tl.constexpr = BLOCK_K // SG * SG
    b_scale_block_n = pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)
    b_scale_k_offs = tl.arange(0, SHUFFLED_SCALE_K)
    scale_k_start_shuffled = scale_k_start * SG
    b_scale_ptrs = b_scale_ptr + b_scale_block_n[:, None] * stride_bsn + (scale_k_start_shuffled + b_scale_k_offs[None, :]) * stride_bsk

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

    num_k_iter = tl.cdiv(k_end - k_start, BLOCK_K // 2)
    for _ in range(0, num_k_iter):
        a = tl.load(a_ptrs)
        b = tl.load(b_ptrs)
        a_scales = tl.load(a_scale_ptrs)

        # B scales: load shuffled, unshuffle in-register (mfma_nonkdim=16 pattern)
        b_scales = tl.load(b_scale_ptrs).reshape(
            BLOCK_N // 32, BLOCK_K // SG // 8, 4, 16, 2, 2, 1,
        ).permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_N, BLOCK_K // SG)

        accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")

        a_ptrs += (BLOCK_K // 2) * stride_ak
        b_ptrs += (BLOCK_K // 2) * stride_bk
        a_scale_ptrs += num_scale_k * stride_ask
        b_scale_ptrs += SHUFFLED_SCALE_K * stride_bsk

    offs_cm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_cn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    c_ptrs = c_ptr + offs_cm[:, None] * stride_cm + offs_cn[None, :] * stride_cn
    c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)

    if SPLIT_K == 1:
        tl.store(c_ptrs, accumulator.to(tl.bfloat16), mask=c_mask)
    else:
        tl.atomic_add(c_ptrs, accumulator, mask=c_mask, sem="relaxed")


def _choose_split_k(M, K_packed, block_k_half=128):
    if M > 32:
        return 1
    max_useful = K_packed // block_k_half
    if max_useful <= 2:
        return 1
    return min(8, max_useful)


def custom_kernel(data: input_t) -> output_t:
    A, B, B_q, B_shuffle, B_scale_sh = data
    M, K = A.shape
    N = B_q.shape[0]
    K_packed = K // 2
    K_scale = K // SCALE_GROUP_SIZE

    from aiter.ops.triton.quant import dynamic_mxfp4_quant

    A_fp4, A_scale_raw = dynamic_mxfp4_quant(A)
    A_q = A_fp4.view(torch.uint8)
    A_scale = A_scale_raw.view(torch.uint8)

    # B scales: reshape AITER's shuffled layout to (N//32, K_scale*32) for
    # in-kernel unshuffle. Free view, no data copy.
    B_scale_raw = B_scale_sh.view(torch.uint8)
    padded_N_scale = B_scale_raw.shape[0]
    padded_K_scale = B_scale_raw.shape[1]
    B_scale_shuffled = B_scale_raw.view(padded_N_scale // 32, padded_K_scale * 32)

    B_q_bytes = B_q.view(torch.uint8)

    BLOCK_M, BLOCK_N, BLOCK_K = 32, 32, 256
    SPLIT_K = _choose_split_k(M, K_packed)

    out_dtype = torch.float32 if SPLIT_K > 1 else torch.bfloat16
    if SPLIT_K > 1:
        C = torch.zeros((M, N), dtype=out_dtype, device=A.device)
    else:
        C = torch.empty((M, N), dtype=out_dtype, device=A.device)

    num_m_tiles = triton.cdiv(M, BLOCK_M)
    num_n_tiles = triton.cdiv(N, BLOCK_N)
    grid = (num_m_tiles * num_n_tiles, SPLIT_K)

    mxfp4_gemm_splitk_kernel[grid](
        A_q, B_q_bytes, C, A_scale, B_scale_shuffled,
        M, N, K_packed,
        A_q.stride(0), A_q.stride(1),
        B_q_bytes.stride(1), B_q_bytes.stride(0),
        C.stride(0), C.stride(1),
        A_scale.stride(0), A_scale.stride(1),
        B_scale_shuffled.stride(0), B_scale_shuffled.stride(1),
        BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
        SPLIT_K=SPLIT_K,
        num_warps=4, num_stages=2,
    )

    if SPLIT_K > 1:
        C = C.to(torch.bfloat16)
    return C
scrolls · 150 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