submission 587635
garrick99 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 208 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-587635?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:226e2edf6bd1c47413dd9bf801d1f758c7c2cb908286695c00823fee734e08b1
license declaredunknown
license concludedunknown
authorsgarrick99
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
MXFP4 GEMM v2 — Triton tl.dot_scaled with native fp4 MFMA.tile-m = 32
BLOCK_M = 32tile-n = 128
BLOCK_N = 128Kernel source
submission.py208 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
MXFP4 GEMM v2 — Triton tl.dot_scaled with native fp4 MFMA.
C[M,N] = A[M,K] @ B[N,K]^T where both A and B are MXFP4 quantized.
Uses tl.dot_scaled with lhs_format='e2m1', rhs_format='e2m1' for native
fp4x4 MFMA on MI355X (gfx950).
Falls back to aiter.gemm_a4w4 if dot_scaled fails.
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
# ============================================================================
# Triton MXFP4 GEMM kernel via dot_scaled
# ============================================================================
@triton.jit
def _mxfp4_gemm_kernel(
A_ptr, B_ptr, # fp4x2 packed: A(M, K//2), B(N, K//2)
A_scale_ptr, B_scale_ptr, # E8M0: A(M_pad, K//32), B(N_pad, K//32)
C_ptr, # bf16: C(M, N)
M, N, K_PACKED, # K_PACKED = K // 2
stride_am, stride_ak,
stride_bn, stride_bk,
stride_asm, stride_ask,
stride_bsn, stride_bsk,
stride_cm, stride_cn,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K_PACKED: tl.constexpr, # packed bytes per K tile (power of 2)
BLOCK_K_SCALE: tl.constexpr, # scale blocks per K tile (BLOCK_K_PACKED // 16)
):
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)
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
m_mask = offs_m < M
n_mask = offs_n < N
for k_start in range(0, K_PACKED, BLOCK_K_PACKED):
offs_k = k_start + tl.arange(0, BLOCK_K_PACKED)
k_mask = offs_k < K_PACKED
# Load A tile: (BLOCK_M, BLOCK_K_PACKED) fp4x2
a = tl.load(
A_ptr + offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak,
mask=m_mask[:, None] & k_mask[None, :], other=0)
# Load B^T tile: (BLOCK_K_PACKED, BLOCK_N) fp4x2
# B stored as (N, K//2), we load transposed
b_t = tl.load(
B_ptr + offs_k[:, None] * stride_bk + offs_n[None, :] * stride_bn,
mask=k_mask[:, None] & n_mask[None, :], other=0)
# Load A scale: (BLOCK_M, BLOCK_K_SCALE) E8M0
k_scale_start = k_start // 16 # 16 packed bytes per scale block
offs_ks = k_scale_start + tl.arange(0, BLOCK_K_SCALE)
a_scale = tl.load(
A_scale_ptr + offs_m[:, None] * stride_asm + offs_ks[None, :] * stride_ask,
mask=m_mask[:, None], other=0)
# Load B scale: (BLOCK_N, BLOCK_K_SCALE) E8M0 — natural (N, K//32) layout
b_scale = tl.load(
B_scale_ptr + offs_n[:, None] * stride_bsn + offs_ks[None, :] * stride_bsk,
mask=n_mask[:, None], other=0)
# dot_scaled: A_fp4 @ B^T_fp4 with block scaling
acc = tl.dot_scaled(
lhs=a,
rhs=b_t,
lhs_scale=a_scale,
rhs_scale=b_scale,
lhs_format='e2m1',
rhs_format='e2m1',
acc=acc,
)
# Store C
c = acc.to(tl.bfloat16)
tl.store(
C_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn, c,
mask=m_mask[:, None] & n_mask[None, :])
# ============================================================================
# Aiter fallback
# ============================================================================
def _aiter_gemm(data):
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
A_fp4, A_bs = dynamic_mxfp4_quant(A)
A_q = A_fp4.view(dtypes.fp4x2)
A_scale_sh = e8m0_shuffle(A_bs).view(dtypes.fp8_e8m0)
return aiter.gemm_a4w4(
A_q, B_shuffle, A_scale_sh, B_scale_sh,
dtype=dtypes.bf16, bpreshuffle=True)
# ============================================================================
# Triton path
# ============================================================================
def _triton_gemm(data):
from aiter.ops.triton.quant import dynamic_mxfp4_quant
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
B = B.contiguous()
m, k = A.shape
n, _ = B.shape
k_packed = k // 2
# Quantize A (unshuffled)
A_fp4, A_scale = dynamic_mxfp4_quant(A)
# Re-quantize B for unshuffled scale (B_q is already unshuffled fp4x2)
_, B_scale = dynamic_mxfp4_quant(B)
# View as uint8 for Triton
a_data = A_fp4
if a_data.dtype != torch.uint8:
a_data = a_data.view(torch.uint8)
b_data = B_q
if not isinstance(b_data, torch.Tensor):
b_data = b_data
if b_data.dtype != torch.uint8:
b_data = b_data.view(torch.uint8)
a_sc = A_scale
if a_sc.dtype != torch.uint8:
a_sc = a_sc.view(torch.uint8)
b_sc = B_scale
if b_sc.dtype != torch.uint8:
b_sc = b_sc.view(torch.uint8)
C = torch.empty((m, n), dtype=torch.bfloat16, device='cuda')
BLOCK_M = 32
BLOCK_N = 128
BLOCK_K_PACKED = 64 # 64 packed bytes = 128 logical fp4 elements
BLOCK_K_SCALE = BLOCK_K_PACKED // 16 # 4 scale blocks
grid = (triton.cdiv(m, BLOCK_M), triton.cdiv(n, BLOCK_N))
_mxfp4_gemm_kernel[grid](
a_data, b_data,
a_sc, b_sc,
C,
m, n, k_packed,
a_data.stride(0), a_data.stride(1),
b_data.stride(0), b_data.stride(1),
a_sc.stride(0), a_sc.stride(1),
b_sc.stride(0), b_sc.stride(1),
C.stride(0), C.stride(1),
BLOCK_M=BLOCK_M,
BLOCK_N=BLOCK_N,
BLOCK_K_PACKED=BLOCK_K_PACKED,
BLOCK_K_SCALE=BLOCK_K_SCALE,
)
return C
# ============================================================================
# Entry point with auto-fallback
# ============================================================================
_use_triton = None
def custom_kernel(data: input_t) -> output_t:
global _use_triton
if _use_triton is None:
try:
result = _triton_gemm(data)
_use_triton = True
return result
except Exception as e:
import sys, traceback
print(f"Triton GEMM FAILED: {type(e).__name__}: {e}", file=sys.stderr)
traceback.print_exc(file=sys.stderr)
_use_triton = False
return _aiter_gemm(data)
if _use_triton:
# Triton dot_scaled wins for small K, aiter wins for large K
A, B, B_q, B_shuffle, B_scale_sh = data
k = A.shape[1]
if k <= 512:
return _triton_gemm(data)
else:
return _aiter_gemm(data)
else:
return _aiter_gemm(data)
scrolls · 208 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