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
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 = 4
num_warps=4, num_stages=2,split-k
"""MXFP4 GEMM — custom Triton kernel with in-kernel scale unshuffle and split-K."""stages = 2
num_warps=4, num_stages=2,tile-k = 32
BLOCK_M, BLOCK_N, BLOCK_K = 32, 32, 256Kernel 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