Skip to content
KernelIndex
Search⌘K

submission 524166

hashkanna · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:aa126ebfd43c8e95dbf9ff459627ebe6d612c79ab142b1e010ccf72ff07b2368
license declaredunknown
license concludedunknown
authorshashkanna
imported2026-08-26

Techniques

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

fp4MXFP4 GEMM for MI355X.

Kernel source

submission.py243 lines
"""
MXFP4 GEMM for MI355X.

Two paths:
  AITER: gemm_a4w4 CK kernel — proven baseline (default)
  TRITON: block_scaled_matmul_kernel_cdna4 — adapted from Triton tutorial, enables
          shape-specific tuning and potential for fused A-quant

Toggle per shape via TRITON_SHAPES dict. Default is aiter for all shapes.
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
import aiter
from aiter import QuantType, dtypes
from aiter.ops.gemm_op_a4w4 import gemm_a4w4_asm

# ─── Module-level init (runs once at import) ──────────────────────────
_quant = aiter.get_triton_quant(QuantType.per_1x32)

# Shapes to route through gemm_a4w4_asm with explicit kernel selection + k-split
# For shapes where the default dispatcher picks a suboptimal kernel
ASM_OVERRIDE_SHAPES: dict[tuple, tuple] = {
    # (M, N, K): (kernelName, log2_k_split)
    # M=16 pads to 32; 32x128 tile with 8-way K-split (log2=3) for K=7168
    # Small tile → 17 tiles * 8 k-split = 136 wave-groups for better CU occupancy
    (16, 2112, 7168): ("_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E", 3),
}

# Shapes to use custom Triton kernel for (manual tl.dot_scaled)
# DISABLED: fp4x2 dtype handling broken with current Triton — needs local GPU debugging
TRITON_SHAPES: dict[tuple, bool] = {}

# Cache for B-side quant in AITER Triton path
_aiter_triton_b_cache: dict = {}  # (M,N,K) -> (B_q_raw, B_scale_raw)

# ─── Scale shuffling for CDNA4 MFMA ──────────────────────────────────
def _shuffle_scales(scales: torch.Tensor, mfma_nonkdim: int) -> torch.Tensor:
    """Shuffle raw E8M0 scales [rows, K//32] → [rows//32, K] for MFMA_SCALE."""
    sm, sn = scales.shape
    if mfma_nonkdim == 32:
        s = scales.view(sm // 32, 32, sn // 8, 4, 2, 1)
        s = s.permute(0, 2, 4, 1, 3, 5).contiguous()
    else:  # 16
        s = scales.view(sm // 32, 2, 16, sn // 8, 2, 4, 1)
        s = s.permute(0, 3, 5, 2, 4, 1, 6).contiguous()
    return s.view(sm // 32, sn * 32)


# ─── Triton CDNA4 MXFP4 GEMM kernel ─────────────────────────────────
# Adapted from triton-lang/triton tutorial 10-block-scaled-matmul.py
# Computes C[M,N] = (A_fp4 * A_scale) @ (B_fp4 * B_scale)^T
# A: [M, K//2] packed fp4x2, B: [K//2, N] packed fp4x2 (transposed)
# A_scales, B_scales: pre-shuffled via _shuffle_scales

@triton.jit
def _mxfp4_gemm_cdna4(
    a_ptr, b_ptr, c_ptr, a_scales_ptr, b_scales_ptr,
    M, N, K,
    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,
    mfma_nonkdim: tl.constexpr,
):
    SCALE_GROUP_SIZE: tl.constexpr = 32

    pid = tl.program_id(axis=0)
    num_pid_n = tl.cdiv(N, BLOCK_N)
    pid_m = pid // num_pid_n
    pid_n = pid % num_pid_n

    # Data pointers (fp4x2 packed: K//2 bytes per row)
    offs_k = tl.arange(0, BLOCK_K // 2)
    offs_am = (pid_m * BLOCK_M + tl.arange(0, BLOCK_M)) % M
    offs_bn = (pid_n * BLOCK_N + tl.arange(0, BLOCK_N)) % N
    a_ptrs = a_ptr + (offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak)
    b_ptrs = b_ptr + (offs_k[:, None] * stride_bk + offs_bn[None, :] * stride_bn)

    # Scale pointers (shuffled: [rows//32, K_scale*32])
    offs_asm = (pid_m * (BLOCK_M // 32) + tl.arange(0, BLOCK_M // 32)) % M
    offs_bsn = (pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)) % N
    offs_ks = tl.arange(0, BLOCK_K // SCALE_GROUP_SIZE * 32)
    a_scale_ptrs = a_scales_ptr + offs_asm[:, None] * stride_asm + offs_ks[None, :] * stride_ask
    b_scale_ptrs = b_scales_ptr + offs_bsn[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk

    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
    num_k_iter = tl.cdiv(K, BLOCK_K // 2)

    for _ in range(0, num_k_iter):
        # Undo scale shuffle in registers
        if mfma_nonkdim == 32:
            a_scales = tl.load(a_scale_ptrs).reshape(
                BLOCK_M // 32, BLOCK_K // SCALE_GROUP_SIZE // 8, 2, 32, 4, 1
            ).permute(0, 3, 1, 4, 2, 5).reshape(BLOCK_M, BLOCK_K // SCALE_GROUP_SIZE)
            b_scales = tl.load(b_scale_ptrs).reshape(
                BLOCK_N // 32, BLOCK_K // SCALE_GROUP_SIZE // 8, 2, 32, 4, 1
            ).permute(0, 3, 1, 4, 2, 5).reshape(BLOCK_N, BLOCK_K // SCALE_GROUP_SIZE)
        elif mfma_nonkdim == 16:
            a_scales = tl.load(a_scale_ptrs).reshape(
                BLOCK_M // 32, BLOCK_K // SCALE_GROUP_SIZE // 8, 4, 16, 2, 2, 1
            ).permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_M, BLOCK_K // SCALE_GROUP_SIZE)
            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)

        a = tl.load(a_ptrs)
        b = tl.load(b_ptrs)
        acc += 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 += BLOCK_K * stride_ask
        b_scale_ptrs += BLOCK_K * stride_bsk

    # Store with write-through cache modifier
    c = acc.to(tl.bfloat16)
    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_ptrs = c_ptr + offs_cm[:, None] * stride_cm + offs_cn[None, :] * stride_cn
    tl.store(c_ptrs, c, mask=(offs_cm[:, None] < M) & (offs_cn[None, :] < N),
             cache_modifier=".wt")


# ─── Triton kernel configs per shape ──────────────────────────────────
# Shape (M, N, K) → (BLOCK_M, BLOCK_N, BLOCK_K, mfma_nonkdim, num_warps, num_stages)
_TRITON_CONFIGS = {
    (4, 2880, 512):    (32, 128, 256, 16, 4, 2),
    (16, 2112, 7168):  (32, 128, 256, 16, 8, 2),
    (32, 4096, 512):   (32, 128, 256, 16, 8, 2),
    (32, 2880, 512):   (32, 128, 256, 16, 8, 2),
    (64, 7168, 2048):  (128, 128, 256, 32, 8, 2),
    (256, 3072, 1536): (128, 128, 256, 32, 8, 2),
}
_DEFAULT_TRITON_CONFIG = (128, 128, 256, 16, 8, 2)


def _pad_to_multiple(val, mult):
    return ((val + mult - 1) // mult) * mult


# Cache B-side prep work (quant + pad + transpose + scale shuffle) — B is constant per shape
_b_cache: dict = {}  # (M,N,K) -> (B_t, B_scale_sh)


def _triton_gemm(A_q, A_scale_raw, B_bf16, M, N, K):
    """Run Triton MXFP4 GEMM. A_q [M,K//2] packed fp4x2, B_bf16 [N,K] bf16."""
    shape_key = (M, N, K)
    cfg = _TRITON_CONFIGS.get(shape_key, _DEFAULT_TRITON_CONFIG)
    BLOCK_M, BLOCK_N, BLOCK_K, mfma_nonkdim, num_warps, num_stages = cfg

    M_pad = _pad_to_multiple(M, BLOCK_M)
    N_pad = _pad_to_multiple(N, BLOCK_N)

    # Cache B-side work (B doesn't change between calls for same shape)
    if shape_key not in _b_cache:
        B_q, B_scale_raw = _quant(B_bf16, shuffle=False)
        # B_q is [N, K//2], B_scale_raw may be [N_padded, K//32]
        if N_pad != N:
            B_q_padded = torch.empty(N_pad, K // 2, dtype=B_q.dtype, device=B_q.device)
            B_q_padded[:N] = B_q
            B_q = B_q_padded
            B_scale_padded = torch.empty(N_pad, K // 32, dtype=B_scale_raw.dtype, device=B_scale_raw.device)
            B_scale_padded[:N] = B_scale_raw[:N]
            B_scale_raw = B_scale_padded
        else:
            B_scale_raw = B_scale_raw[:N]
        B_t = B_q.view(torch.uint8).T.contiguous()
        B_scale_sh = _shuffle_scales(B_scale_raw, mfma_nonkdim).view(torch.uint8).contiguous()
        _b_cache[shape_key] = (B_t, B_scale_sh)

    B_t, B_scale_sh = _b_cache[shape_key]

    # A-side: pad if needed + shuffle scales (A changes every call)
    if M_pad != M:
        A_q_padded = torch.empty(M_pad, K // 2, dtype=A_q.dtype, device=A_q.device)
        A_q_padded[:M] = A_q
        A_q = A_q_padded
        A_scale_padded = torch.empty(M_pad, K // 32, dtype=A_scale_raw.dtype, device=A_scale_raw.device)
        A_scale_padded[:M] = A_scale_raw[:M]
        A_scale_raw = A_scale_padded
    else:
        A_scale_raw = A_scale_raw[:M]

    A_scale_sh = _shuffle_scales(A_scale_raw, mfma_nonkdim)

    # Output
    C = torch.empty(M_pad, N_pad, dtype=torch.bfloat16, device="cuda")

    grid = (triton.cdiv(M_pad, BLOCK_M) * triton.cdiv(N_pad, BLOCK_N), 1)

    # View fp4x2/e8m0 as uint8 — Triton can't handle fp4x2 dtype directly
    A_q_u8 = A_q.view(torch.uint8)
    A_sc_u8 = A_scale_sh.view(torch.uint8)
    # B_t and B_scale_sh are already uint8 from cache

    _mxfp4_gemm_cdna4[grid](
        A_q_u8, B_t, C, A_sc_u8, B_scale_sh,
        M_pad, N_pad, K,
        A_q_u8.stride(0), A_q_u8.stride(1),
        B_t.stride(0), B_t.stride(1),
        C.stride(0), C.stride(1),
        A_sc_u8.stride(0), A_sc_u8.stride(1),
        B_scale_sh.stride(0), B_scale_sh.stride(1),
        BLOCK_M, BLOCK_N, BLOCK_K, mfma_nonkdim,
        num_warps=num_warps, num_stages=num_stages,
        matrix_instr_nonkdim=mfma_nonkdim,
    )
    return C[:M, :N]


# ─── Entry point ──────────────────────────────────────────────────────
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]

    asm_cfg = ASM_OVERRIDE_SHAPES.get((M, N, K))
    if asm_cfg is not None:
        # Direct ASM path: call gemm_a4w4_asm with explicit kernel + k-split
        kernel_name, log2_k_split = asm_cfg
        A_q, A_scale_sh = _quant(A, shuffle=True)
        M_pad = ((M + 31) // 32) * 32
        out = torch.empty(M_pad, N, dtype=torch.bfloat16, device="cuda")
        gemm_a4w4_asm(A_q, B_shuffle, A_scale_sh, B_scale_sh, out,
                       kernel_name, None, 1.0, 0.0, True, log2_k_split)
        return out[:M]

    if TRITON_SHAPES.get((M, N, K), False):
        # Custom Triton path (disabled): uses manual tl.dot_scaled kernel
        A_q, A_scale_raw = _quant(A, shuffle=False)
        return _triton_gemm(A_q, A_scale_raw, B, M, N, K)

    # Aiter ASM/CK path (default): proven correct, matches reference
    A_q, A_scale_sh = _quant(A, shuffle=True)
    return aiter.gemm_a4w4(
        A_q, B_shuffle, A_scale_sh, B_scale_sh,
        dtype=dtypes.bf16, bpreshuffle=True,
    )
scrolls · 243 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