submission 587049
gau.nernst · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 331 lines, June 9 Researcher Reciprocity License v1.0.
submission_v2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-587049?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:b0af054b38c912c6bfa67eba90c0bf01c6391270ffb75277a2626270a4537498
license declaredunknown
license concludedunknown
authorsgau.nernst
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
submission_v2.py331 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
import torch
import triton
import triton.language as tl
from task import input_t, output_t
from torch import Tensor
@triton.jit
def _dynamic_mxfp4_quant_kernel(
x_ptr,
x_fp4_ptr,
bs_ptr,
stride_xm,
stride_xn,
M,
N: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
NUM_ITER: tl.constexpr,
NUM_STAGES: tl.constexpr,
EVEN_MN: tl.constexpr,
):
pid_m = tl.program_id(0)
start_n = tl.program_id(1) * NUM_ITER
BLOCK_SF_N: tl.constexpr = BLOCK_N // 32
for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):
x_offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
x_offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
x_offs = x_offs_m[:, None] * stride_xm + x_offs_n[None, :] * stride_xn
if EVEN_MN:
x = tl.load(x_ptr + x_offs, cache_modifier=".cg").to(tl.float32)
else:
x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
x = tl.load(x_ptr + x_offs, mask=x_mask, cache_modifier=".cg").to(tl.float32)
x = x.reshape(BLOCK_M, BLOCK_SF_N, 32)
# Calculate scale
amax = tl.max(tl.abs(x), axis=-1, keep_dims=True)
amax = amax.to(tl.int32, bitcast=True)
amax = (amax + 0x200000).to(tl.uint32, bitcast=True) & 0xFF800000
amax = amax.to(tl.float32, bitcast=True)
scale_e8m0_unbiased = tl.log2(amax).floor() - 2
scale_e8m0_unbiased = tl.clamp(scale_e8m0_unbiased, min=-127, max=127)
# blockscale_e8m0
bs_e8m0 = scale_e8m0_unbiased.to(tl.uint8) + 127 # in fp32, we have 2^(e - 127)
bs_e8m0 = bs_e8m0.reshape(BLOCK_M, BLOCK_SF_N)
quant_scale = tl.exp2(-scale_e8m0_unbiased)
# Compute quantized x
qx = x * quant_scale
qx_lo, qx_hi = qx.reshape(BLOCK_M, BLOCK_N // 2, 2).split()
ones = tl.full((BLOCK_M, BLOCK_N // 2), 1.0, dtype=tl.float32)
x_fp4 = tl.inline_asm_elementwise(
"V_CVT_SCALEF32_PK_FP4_F32 $0, $1, $2, $3;",
constraints="=v,v,v,v",
args=[qx_lo, qx_hi, ones],
dtype=tl.uint32, # we have to use 32-bit for v constraint?
is_pure=True,
pack=1,
).to(tl.uint8)
# store x_fp4
out_offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
out_offs_n = pid_n * BLOCK_N // 2 + tl.arange(0, BLOCK_N // 2)
out_offs = out_offs_m[:, None] * (N // 2) + out_offs_n[None, :]
if EVEN_MN:
tl.store(x_fp4_ptr + out_offs, x_fp4)
else:
out_mask = (out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :]
tl.store(x_fp4_ptr + out_offs, x_fp4, mask=out_mask)
# store blockscale
bs_offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
bs_offs_n = pid_n * BLOCK_SF_N + tl.arange(0, BLOCK_SF_N)
# https://github.com/ROCm/aiter/blob/dc7fefcd/aiter/ops/triton/_triton_kernels/quant/fused_mxfp4_quant.py#L173
# .view(M/32, 2, 16, N/8, 2, 4) | [m0, m1, m2, n0, n1, n2]
# .permute(0, 3, 5, 2, 4, 1)
# => [M/32, N/8, 4, 16, 2, 2] | [m0, n0, n2, m2, n1, m1]
SCALE_N_PAD: tl.constexpr = tl.cdiv(N, 256) * 8
bs_offs_0 = bs_offs_m[:, None] // 32
bs_offs_1 = bs_offs_m[:, None] % 32
bs_offs_2 = bs_offs_1 % 16
bs_offs_1 = bs_offs_1 // 16
bs_offs_3 = bs_offs_n[None, :] // 8
bs_offs_4 = bs_offs_n[None, :] % 8
bs_offs_5 = bs_offs_4 % 4
bs_offs_4 = bs_offs_4 // 4
bs_offs = (
bs_offs_1
+ bs_offs_4 * 2
+ bs_offs_2 * 2 * 2
+ bs_offs_5 * 2 * 2 * 16
+ bs_offs_3 * 2 * 2 * 16 * 4
+ bs_offs_0 * 2 * 16 * SCALE_N_PAD
)
# bs_offs = bs_offs_m[:, None] * stride_bs_m + bs_offs_n[None, :] * stride_bs_n
if EVEN_MN:
tl.store(bs_ptr + bs_offs, bs_e8m0)
else:
SCALE_N: tl.constexpr = N // 32
bs_mask = (bs_offs_m < M)[:, None] & (bs_offs_n < SCALE_N)[None, :]
tl.store(bs_ptr + bs_offs, bs_e8m0, mask=bs_mask)
def dynamic_mxfp4_quant(x: Tensor) -> tuple[Tensor, Tensor]:
M, N = x.shape
# M_pad = triton.cdiv(M, 256) * 256
M_pad = triton.cdiv(M, 32) * 32
# x_fp4 = x.new_empty((M_pad, N // 2), dtype=torch.uint8)
x_fp4 = x.new_empty((M, N // 2), dtype=torch.uint8)
bs_shape = (M_pad, triton.cdiv(N, 256) * 8)
blockscale_e8m0 = x.new_empty(bs_shape, dtype=torch.uint8)
# for large N values
if M <= 32:
NUM_ITER = 1
BLOCK_M = triton.next_power_of_2(M)
BLOCK_N = 32
NUM_WARPS = 1
NUM_STAGES = 1
else:
NUM_ITER = 4
BLOCK_M = 64
BLOCK_N = 64
NUM_WARPS = 4
NUM_STAGES = 2
if N <= 16384:
BLOCK_M = 32
BLOCK_N = 128
# for small N values
if N <= 1024:
NUM_ITER = 1
NUM_STAGES = 1
NUM_WARPS = 4
BLOCK_N = min(256, triton.next_power_of_2(N))
# BLOCK_N needs to be multiple of 32
BLOCK_N = max(32, BLOCK_N)
BLOCK_M = min(8, triton.next_power_of_2(M))
EVEN_MN = (M % BLOCK_M == 0) and (N % (BLOCK_N * NUM_ITER) == 0)
grid = (triton.cdiv(M, BLOCK_M), triton.cdiv(N, BLOCK_N * NUM_ITER))
_dynamic_mxfp4_quant_kernel[grid](
x,
x_fp4,
blockscale_e8m0,
*x.stride(),
M=M,
N=N,
NUM_ITER=NUM_ITER,
BLOCK_M=BLOCK_M,
BLOCK_N=BLOCK_N,
NUM_STAGES=NUM_STAGES,
EVEN_MN=EVEN_MN,
num_warps=NUM_WARPS,
waves_per_eu=0,
num_stages=1,
)
x_fp4 = x_fp4.view(torch.float4_e2m1fn_x2)
blockscale_e8m0 = blockscale_e8m0.view(torch.float8_e8m0fnu)
return x_fp4, blockscale_e8m0
# https://triton-lang.org/main/getting-started/tutorials/10-block-scaled-matmul.html
@triton.jit
def mxfp4_mm_kernel(
A_ptr,
B_ptr,
C_ptr,
SFA_ptr,
SFB_ptr,
M,
N: tl.constexpr,
K: tl.constexpr,
stride_am,
stride_bn,
# Meta-parameters
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
A_FIRST: tl.constexpr,
):
tl.static_assert(K % BLOCK_K == 0)
pid = tl.program_id(axis=0)
# TODO: XCD remapping
num_pid_n = tl.cdiv(N, BLOCK_N)
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
offs_k = tl.arange(0, BLOCK_K // 2)
offs_k_split = offs_k
offs_am = (pid_m * BLOCK_M + tl.arange(0, BLOCK_M)) % M
offs_bn = (pid_n * BLOCK_N + tl.arange(0, BLOCK_N)) % N
if A_FIRST:
A_ptrs = A_ptr + (offs_am[:, None] * stride_am + offs_k_split[None, :])
B_ptrs = B_ptr + (offs_k_split[:, None] + offs_bn[None, :] * stride_bn)
else:
A_ptrs = A_ptr + (offs_am[None, :] * stride_am + offs_k_split[:, None])
B_ptrs = B_ptr + (offs_k_split[None, :] + offs_bn[:, None] * stride_bn)
# Create pointers for the first block of A and B scales
offs_asm = pid_m * (BLOCK_M // 32) + tl.arange(0, BLOCK_M // 32)
offs_bsn = pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)
offs_ks = tl.arange(0, BLOCK_K)
# SFA/SFB are packed as [M/32, K/32, 256]
SFA_ptrs = SFA_ptr + (offs_asm[:, None] * K + offs_ks[None, :])
SFB_ptrs = SFB_ptr + (offs_bsn[:, None] * K + offs_ks[None, :])
if A_FIRST:
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
else:
acc = tl.zeros((BLOCK_N, BLOCK_M), dtype=tl.float32)
for _ in range(K // BLOCK_K):
# load B first doesn't seem to be faster
# TODO: don't shuffle SFA so that we can use smaller BLOCK_M?
# "undo" SF shuffle
SFA = (
tl.load(SFA_ptrs)
.reshape(BLOCK_M // 32, BLOCK_K // 256, 4, 16, 2, 2)
.permute(0, 5, 3, 1, 4, 2)
.reshape(BLOCK_M, BLOCK_K // 32)
)
SFB = (
tl.load(SFB_ptrs, cache_modifier=".cg") # doesn't seem to matter much here
.reshape(BLOCK_N // 32, BLOCK_K // 256, 4, 16, 2, 2)
.permute(0, 5, 3, 1, 4, 2)
.reshape(BLOCK_N, BLOCK_K // 32)
)
A = tl.load(A_ptrs)
B = tl.load(B_ptrs, cache_modifier=".cg")
if A_FIRST:
acc = tl.dot_scaled(A, SFA, "e2m1", B, SFB, "e2m1", acc=acc)
else:
acc = tl.dot_scaled(B, SFB, "e2m1", A, SFA, "e2m1", acc=acc)
A_ptrs += BLOCK_K // 2
B_ptrs += BLOCK_K // 2
SFA_ptrs += BLOCK_K
SFB_ptrs += BLOCK_K
if A_FIRST:
offs_cm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)[:, None]
offs_cn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)[None, :]
else:
offs_cm = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)[None, :]
offs_cn = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)[:, None]
c_ptrs = C_ptr + (offs_cm * N + offs_cn)
if N % BLOCK_N == 0:
c_mask = offs_cm < M
else:
c_mask = (offs_cm < M) & (offs_cn < N)
tl.store(c_ptrs, acc, mask=c_mask, cache_modifier=".wt")
# (M, N, K): (BM, BN, BK, num_stages, num_warps)
config_map = {
(4, 2880, 512): (32, 32, 512, 1, 4, True),
(16, 2112, 7168): (32, 32, 1024, 4, 4, False),
(32, 4096, 512): (32, 32, 512, 1, 4, True),
(32, 2880, 512): (32, 32, 512, 1, 4, True),
(64, 7168, 2048): (64, 32, 512, 4, 4, True),
(256, 3072, 1536): (128, 32, 512, 3, 4, True),
}
default_config = (64, 64, 512, 4, 4, True)
def mxfp4_mm(A: Tensor, B: Tensor, SFA: Tensor, SFB: Tensor):
M = A.shape[0]
N = B.shape[0]
K = A.shape[1] * 2
out = torch.empty((M, N), device=A.device, dtype=torch.bfloat16)
BLOCK_M, BLOCK_N, BLOCK_K, num_stages, num_warps, A_FIRST = config_map.get((M, N, K), default_config)
grid = (triton.cdiv(M, BLOCK_M) * triton.cdiv(N, BLOCK_N), 1)
mxfp4_mm_kernel[grid](
A.view(torch.uint8),
B.view(torch.uint8),
out,
SFA.view(torch.uint8),
SFB.view(torch.uint8),
M,
N,
K,
A.stride(0),
B.stride(0),
BLOCK_M,
BLOCK_N,
BLOCK_K,
A_FIRST,
num_warps=num_warps,
num_stages=num_stages,
matrix_instr_nonkdim=16,
)
return out
def custom_kernel(data: input_t) -> output_t:
A, _, Bq, Bq_shfl, SFB = data
Aq, SFA = dynamic_mxfp4_quant(A)
out_gemm = mxfp4_mm(Aq, Bq, SFA, SFB)
return out_gemm
scrolls · 331 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