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
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.
fp4
MXFP4 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