submission 617263
mega-dmitriy · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 391 lines, June 9 Researcher Reciprocity License v1.0.
submission_v469_pro2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-617263?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:ebafd308fc625fca71268e0d0d48911e5263a8d8d4d217aaeb2ac5244f448864
license declaredunknown
license concludedunknown
authorsmega-dmitriy
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 4
num_warps=4, num_stages=2, matrix_instr_nonkdim=16,split-k
def _fused_splitk_kernel(stages = 2
SCALE_GP=16, GROUP_SIZE_M=8, NUM_STAGES=2,tile-k = 16
BM, BN, BK = 16, 64, 256tile-n = 128
RED_BN = 128Kernel source
submission_v469_pro2.py391 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
V469_pro2: Micro-optimizations on v466 for ~1% geomean improvement.
Changes:
1. Eliminate scale_bc = scale_f32 + tl.zeros() broadcast — use reshape trick
2. Pre-compute N-dependent scale shuffle terms outside K loop
3. BF16 amax: compute amax in bf16, only convert per-group max to f32
4. .wt cache modifier on quant intermediate stores (K=1536)
"""
import torch
import triton
import triton.language as tl
from task import input_t, output_t
import aiter
_ASM_KERNEL = "_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E"
_cache = {}
@triton.jit
def _quant_shuffle_kernel(
a_ptr, fp4_out_ptr, scale_out_ptr,
stride_am, stride_ak,
K_HALF: tl.constexpr, scaleN_pad: tl.constexpr,
M_VAL: tl.constexpr, BLOCK_K: tl.constexpr,
):
m = tl.program_id(0)
blk_k = tl.program_id(1)
if m >= M_VAL:
return
GPB: tl.constexpr = BLOCK_K // 32
k_start = blk_k * BLOCK_K
offs_k = tl.arange(0, BLOCK_K)
# BF16 amax: skip bulk f32 conversion
a_bf16 = tl.load(a_ptr + m * stride_am + (k_start + offs_k) * stride_ak)
a_grouped_bf16 = tl.reshape(a_bf16, [GPB, 32])
amax_bf16 = tl.max(tl.abs(a_grouped_bf16), axis=1, keep_dims=True)
amax_f32 = amax_bf16.to(tl.float32) # only GPB values converted
amax_u32 = amax_f32.to(tl.uint32, bitcast=True)
amax_u32 = (amax_u32 + 0x200000) & 0xFF800000
amax_exp = (amax_u32 >> 23).to(tl.int32)
scale_exp = tl.minimum(tl.maximum(amax_exp - 2, 0), 254)
bs_e8m0 = scale_exp.to(tl.uint8)
scale_f32 = (scale_exp.to(tl.uint32) << 23).to(tl.float32, bitcast=True)
# Still need f32 for HW FP4 convert inputs
a_f32 = a_bf16.to(tl.float32)
a_grouped = tl.reshape(a_f32, [GPB, 32])
a_pairs = tl.reshape(a_grouped, [GPB, 16, 2])
a_even, a_odd = tl.split(a_pairs)
# Scale broadcast via reshape: [GPB,1] -> repeat 16x -> [GPB,16] -> flatten
scale_1d = tl.reshape(scale_f32, [GPB])
scale_rep = scale_1d[:, None] * tl.full([1, 16], 1.0, dtype=tl.float32)
a_even_flat = tl.reshape(a_even, [GPB * 16])
a_odd_flat = tl.reshape(a_odd, [GPB * 16])
scale_flat = tl.reshape(scale_rep, [GPB * 16])
packed_u32 = tl.inline_asm_elementwise(
asm="v_cvt_scalef32_pk_fp4_f32 $0, $1, $2, $3",
constraints="=&v,v,v,v",
args=[a_even_flat, a_odd_flat, scale_flat],
dtype=tl.uint32, is_pure=True, pack=1,
)
packed_u8 = (packed_u32 & 0xFF).to(tl.uint8)
fp4_flat = tl.reshape(packed_u8, [GPB * 16])
fp4_offs = tl.arange(0, GPB * 16)
tl.store(fp4_out_ptr + m * K_HALF + k_start // 2 + fp4_offs, fp4_flat, cache_modifier=".wt")
g_base = k_start // 32
g_vals = g_base + tl.arange(0, GPB)
sh_off = ((m % 32 // 16) + (g_vals % 8 // 4) * 2 + (m % 16) * 4
+ (g_vals % 4) * 64 + (g_vals // 8) * 256
+ (m // 32) * (32 * scaleN_pad))
scale_bytes = tl.reshape(bs_e8m0, [GPB])
tl.store(scale_out_ptr + sh_off, scale_bytes, cache_modifier=".wt")
@triton.jit
def _fused_kernel(
a_ptr, b_ptr, c_ptr, b_scales_ptr,
N,
stride_am, stride_ak, stride_bn, stride_cm, stride_cn,
scaleN_pad_B,
M_VAL: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr, K_HALF: tl.constexpr, SCALE_GP: tl.constexpr,
GROUP_SIZE_M: tl.constexpr, NUM_STAGES: tl.constexpr,
MASK_M: tl.constexpr, MASK_N: tl.constexpr,
):
pid = tl.program_id(0)
NUM_PID_M: tl.constexpr = (M_VAL + BLOCK_M - 1) // BLOCK_M
num_pid_n = tl.cdiv(N, BLOCK_N)
num_pid_in_group = GROUP_SIZE_M * num_pid_n
group_id = pid // num_pid_in_group
first_pid_m = group_id * GROUP_SIZE_M
group_size_m = min(NUM_PID_M - first_pid_m, GROUP_SIZE_M)
pid_m = first_pid_m + ((pid % num_pid_in_group) % group_size_m)
pid_n = (pid % num_pid_in_group) // group_size_m
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
offs_k = tl.arange(0, BLOCK_K)
NQB: tl.constexpr = BLOCK_K // SCALE_GP
a_bf16_offs_k = tl.arange(0, BLOCK_K * 2)
a_ptrs = a_ptr + offs_m[:, None] * stride_am
b_ptrs = b_ptr + offs_n[None, :] * stride_bn + offs_k[:, None]
offs_qs = tl.arange(0, NQB)
accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
# Pre-compute N-dependent scale shuffle terms (loop-invariant)
n_val = offs_n[:, None]
n_sh_base = (n_val % 32 // 16) + (n_val % 16) * 4 + (n_val // 32) * (32 * scaleN_pad_B)
for k_block in tl.range(0, K_HALF, BLOCK_K, num_stages=NUM_STAGES):
if MASK_M:
a_bf16 = tl.load(a_ptrs + (k_block * 2 + a_bf16_offs_k[None, :]) * stride_ak, mask=(offs_m[:, None] < M_VAL), other=0.0)
else:
a_bf16 = tl.load(a_ptrs + (k_block * 2 + a_bf16_offs_k[None, :]) * stride_ak)
# BF16 amax: only convert per-group max to f32
a_grouped_bf16 = tl.reshape(a_bf16, [BLOCK_M * NQB, 32])
amax_bf16 = tl.max(tl.abs(a_grouped_bf16), axis=1, keep_dims=True)
amax_f32 = amax_bf16.to(tl.float32)
amax_u32 = amax_f32.to(tl.uint32, bitcast=True)
amax_u32 = (amax_u32 + 0x200000) & 0xFF800000
amax_exp = (amax_u32 >> 23).to(tl.int32)
scale_exp = tl.minimum(tl.maximum(amax_exp - 2, 0), 254)
bs_e8m0 = scale_exp.to(tl.uint8)
scale_f32 = (scale_exp.to(tl.uint32) << 23).to(tl.float32, bitcast=True)
# Convert to f32 for HW FP4 (still needed for F32 variant input)
a_f32 = a_bf16.to(tl.float32)
a_grouped = tl.reshape(a_f32, [BLOCK_M * NQB, 32])
a_pairs = tl.reshape(a_grouped, [BLOCK_M * NQB, 16, 2])
a_even, a_odd = tl.split(a_pairs)
# Scale broadcast: reshape [G,1] -> [G] -> [:, None] * ones -> [G,16]
scale_1d = tl.reshape(scale_f32, [BLOCK_M * NQB])
scale_rep = scale_1d[:, None] * tl.full([1, 16], 1.0, dtype=tl.float32)
a_even_flat = tl.reshape(a_even, [BLOCK_M * NQB * 16])
a_odd_flat = tl.reshape(a_odd, [BLOCK_M * NQB * 16])
scale_flat = tl.reshape(scale_rep, [BLOCK_M * NQB * 16])
packed_u32 = tl.inline_asm_elementwise(
asm="v_cvt_scalef32_pk_fp4_f32 $0, $1, $2, $3",
constraints="=&v,v,v,v",
args=[a_even_flat, a_odd_flat, scale_flat],
dtype=tl.uint32, is_pure=True, pack=1,
)
packed_u8 = (packed_u32 & 0xFF).to(tl.uint8)
a_quant = tl.reshape(packed_u8, [BLOCK_M, BLOCK_K])
a_scales = tl.reshape(bs_e8m0, [BLOCK_M, NQB])
if MASK_N:
b = tl.load(b_ptrs + k_block, mask=(offs_n[None, :] < N), other=0)
else:
b = tl.load(b_ptrs + k_block)
# B scale load with pre-computed N terms
g_base = k_block // SCALE_GP
g_val = g_base + offs_qs[None, :]
sh_off = n_sh_base + (g_val % 8 // 4) * 2 + (g_val % 4) * 64 + (g_val // 8) * 256
if MASK_N:
b_sc = tl.load(b_scales_ptr + sh_off, mask=(offs_n[:, None] < N), other=0)
else:
b_sc = tl.load(b_scales_ptr + sh_off)
accumulator += tl.dot_scaled(a_quant, a_scales, "e2m1", b, b_sc, "e2m1")
c = accumulator.to(tl.bfloat16)
c_mask = (offs_m[:, None] < M_VAL) & (offs_n[None, :] < N)
c_ptrs = c_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
tl.store(c_ptrs, c, mask=c_mask, cache_modifier=".wt")
@triton.jit
def _fused_splitk_kernel(
a_ptr, b_ptr, workspace_ptr, b_scales_ptr,
N, stride_am, stride_ak, stride_bn, stride_wm, stride_wn,
scaleN_pad_B,
M_VAL: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr, SCALE_GP: tl.constexpr,
K_PER_SPLIT: tl.constexpr, NUM_KSPLIT: tl.constexpr,
MASK_M: tl.constexpr, MASK_N: tl.constexpr,
):
pid_mn = tl.program_id(0)
pid_k = tl.program_id(1)
num_pid_m = tl.cdiv(M_VAL, BLOCK_M)
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)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
offs_k = tl.arange(0, BLOCK_K)
NQB: tl.constexpr = BLOCK_K // SCALE_GP
a_bf16_offs_k = tl.arange(0, BLOCK_K * 2)
a_ptrs = a_ptr + offs_m[:, None] * stride_am
b_ptrs = b_ptr + offs_n[None, :] * stride_bn + offs_k[:, None]
offs_qs = tl.arange(0, NQB)
k_start = pid_k * K_PER_SPLIT
k_end = k_start + K_PER_SPLIT
accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
# Pre-compute N-dependent scale shuffle terms
n_val = offs_n[:, None]
n_sh_base = (n_val % 32 // 16) + (n_val % 16) * 4 + (n_val // 32) * (32 * scaleN_pad_B)
for k_block in tl.range(k_start, k_end, BLOCK_K):
if MASK_M:
a_bf16 = tl.load(a_ptrs + (k_block * 2 + a_bf16_offs_k[None, :]) * stride_ak, mask=(offs_m[:, None] < M_VAL), other=0.0)
else:
a_bf16 = tl.load(a_ptrs + (k_block * 2 + a_bf16_offs_k[None, :]) * stride_ak)
# BF16 amax
a_grouped_bf16 = tl.reshape(a_bf16, [BLOCK_M * NQB, 32])
amax_bf16 = tl.max(tl.abs(a_grouped_bf16), axis=1, keep_dims=True)
amax_f32 = amax_bf16.to(tl.float32)
amax_u32 = amax_f32.to(tl.uint32, bitcast=True)
amax_u32 = (amax_u32 + 0x200000) & 0xFF800000
amax_exp = (amax_u32 >> 23).to(tl.int32)
scale_exp = tl.minimum(tl.maximum(amax_exp - 2, 0), 254)
bs_e8m0 = scale_exp.to(tl.uint8)
scale_f32 = (scale_exp.to(tl.uint32) << 23).to(tl.float32, bitcast=True)
a_f32 = a_bf16.to(tl.float32)
a_grouped = tl.reshape(a_f32, [BLOCK_M * NQB, 32])
a_pairs = tl.reshape(a_grouped, [BLOCK_M * NQB, 16, 2])
a_even, a_odd = tl.split(a_pairs)
scale_1d = tl.reshape(scale_f32, [BLOCK_M * NQB])
scale_rep = scale_1d[:, None] * tl.full([1, 16], 1.0, dtype=tl.float32)
a_even_flat = tl.reshape(a_even, [BLOCK_M * NQB * 16])
a_odd_flat = tl.reshape(a_odd, [BLOCK_M * NQB * 16])
scale_flat = tl.reshape(scale_rep, [BLOCK_M * NQB * 16])
packed_u32 = tl.inline_asm_elementwise(
asm="v_cvt_scalef32_pk_fp4_f32 $0, $1, $2, $3",
constraints="=&v,v,v,v",
args=[a_even_flat, a_odd_flat, scale_flat],
dtype=tl.uint32, is_pure=True, pack=1,
)
packed_u8 = (packed_u32 & 0xFF).to(tl.uint8)
a_quant = tl.reshape(packed_u8, [BLOCK_M, BLOCK_K])
a_scales = tl.reshape(bs_e8m0, [BLOCK_M, NQB])
if MASK_N:
b = tl.load(b_ptrs + k_block, mask=(offs_n[None, :] < N), other=0)
else:
b = tl.load(b_ptrs + k_block)
g_base = k_block // SCALE_GP
g_val = g_base + offs_qs[None, :]
sh_off = n_sh_base + (g_val % 8 // 4) * 2 + (g_val % 4) * 64 + (g_val // 8) * 256
if MASK_N:
b_sc = tl.load(b_scales_ptr + sh_off, mask=(offs_n[:, None] < N), other=0)
else:
b_sc = tl.load(b_scales_ptr + sh_off)
accumulator += tl.dot_scaled(a_quant, a_scales, "e2m1", b, b_sc, "e2m1")
w_mask = (offs_m[:, None] < M_VAL) & (offs_n[None, :] < N)
w_ptrs = workspace_ptr + pid_k * stride_wm * M_VAL + offs_m[:, None] * stride_wm + offs_n[None, :] * stride_wn
tl.store(w_ptrs, accumulator, mask=w_mask)
@triton.jit
def _reduce_kernel(
workspace_ptr, c_ptr,
M_VAL: tl.constexpr, N,
stride_wm, stride_wn, stride_cm, stride_cn,
NUM_KSPLIT: tl.constexpr, BLOCK_N: tl.constexpr,
):
m = tl.program_id(0)
pid_n = tl.program_id(1)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
n_mask = offs_n < N
acc = tl.zeros((BLOCK_N,), dtype=tl.float32)
for k in range(NUM_KSPLIT):
vals = tl.load(workspace_ptr + k * M_VAL * stride_wm + m * stride_wm + offs_n * stride_wn, mask=n_mask, other=0.0)
acc += vals
tl.store(c_ptr + m * stride_cm + offs_n * stride_cn, acc.to(tl.bfloat16), mask=n_mask)
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
M, K = A.shape
N = B_q.shape[0]
K_half = K // 2
key = (M, N, K)
if key not in _cache:
out = torch.empty((M, N), dtype=torch.bfloat16, device=A.device)
if K >= 4096:
NUM_KSPLIT = K // 512
workspace = torch.empty((NUM_KSPLIT, M, N), dtype=torch.float32, device=A.device)
_cache[key] = (out, workspace, NUM_KSPLIT)
elif K > 512 and K != 2048:
scale_n = K // 32
scale_n_pad = triton.cdiv(scale_n, 8) * 8
m_pad = triton.cdiv(M, 256) * 256
fp4_buf = torch.empty((M, K_half), dtype=torch.uint8, device=A.device)
scale_buf = torch.zeros((m_pad, scale_n_pad), dtype=torch.uint8, device=A.device)
_cache[key] = (out, fp4_buf, scale_buf, scale_n_pad, m_pad)
else:
_cache[key] = (out,)
cached = _cache[key]
out = cached[0]
if K <= 512:
B_q_u8 = B_q.view(torch.uint8)
B_scale_u8 = B_scale_sh.view(torch.uint8)
scaleN_pad_B = triton.cdiv(K // 32, 8) * 8
BM, BN, BK = 16, 64, 256
MASK_M = (M % BM) != 0
MASK_N = (N % BN) != 0
grid = (triton.cdiv(M, BM) * triton.cdiv(N, BN),)
_fused_kernel[grid](
A, B_q_u8, out, B_scale_u8, N,
A.stride(0), A.stride(1), B_q_u8.stride(0),
out.stride(0), out.stride(1), scaleN_pad_B,
M_VAL=M, BLOCK_M=BM, BLOCK_N=BN, BLOCK_K=BK, K_HALF=K_half,
SCALE_GP=16, GROUP_SIZE_M=8, NUM_STAGES=2,
MASK_M=MASK_M, MASK_N=MASK_N,
num_warps=4, num_stages=2, matrix_instr_nonkdim=16,
)
elif K == 2048:
B_q_u8 = B_q.view(torch.uint8)
B_scale_u8 = B_scale_sh.view(torch.uint8)
scaleN_pad_B = triton.cdiv(K // 32, 8) * 8
BM, BN, BK = 16, 128, 256
MASK_M = (M % BM) != 0
MASK_N = (N % BN) != 0
grid = (triton.cdiv(M, BM) * triton.cdiv(N, BN),)
_fused_kernel[grid](
A, B_q_u8, out, B_scale_u8, N,
A.stride(0), A.stride(1), B_q_u8.stride(0),
out.stride(0), out.stride(1), scaleN_pad_B,
M_VAL=M, BLOCK_M=BM, BLOCK_N=BN, BLOCK_K=BK, K_HALF=K_half,
SCALE_GP=16, GROUP_SIZE_M=8, NUM_STAGES=2,
MASK_M=MASK_M, MASK_N=MASK_N,
num_warps=8, num_stages=2, matrix_instr_nonkdim=16,
)
elif K >= 4096:
_, workspace, NUM_KSPLIT = cached
B_q_u8 = B_q.view(torch.uint8)
B_scale_u8 = B_scale_sh.view(torch.uint8)
scaleN_pad_B = triton.cdiv(K // 32, 8) * 8
BM, BN, BK = 16, 128, 256
MASK_M = (M % BM) != 0
MASK_N = (N % BN) != 0
K_PER_SPLIT = K_half // NUM_KSPLIT
num_mn_tiles = triton.cdiv(M, BM) * triton.cdiv(N, BN)
grid_fused = (num_mn_tiles, NUM_KSPLIT)
_fused_splitk_kernel[grid_fused](
A, B_q_u8, workspace, B_scale_u8,
N, A.stride(0), A.stride(1), B_q_u8.stride(0),
workspace.stride(1), workspace.stride(2),
scaleN_pad_B,
M_VAL=M, BLOCK_M=BM, BLOCK_N=BN, BLOCK_K=BK,
SCALE_GP=16, K_PER_SPLIT=K_PER_SPLIT, NUM_KSPLIT=NUM_KSPLIT,
MASK_M=MASK_M, MASK_N=MASK_N,
num_warps=8, num_stages=2, matrix_instr_nonkdim=16,
)
RED_BN = 128
grid_reduce = (M, triton.cdiv(N, RED_BN))
_reduce_kernel[grid_reduce](
workspace, out, M_VAL=M, N=N,
stride_wm=workspace.stride(1), stride_wn=workspace.stride(2),
stride_cm=out.stride(0), stride_cn=out.stride(1),
NUM_KSPLIT=NUM_KSPLIT, BLOCK_N=RED_BN,
num_warps=4,
)
else:
_, fp4_buf, scale_buf, scale_n_pad, m_pad = cached
BLOCK_K = 256
num_k_blocks = K // BLOCK_K
grid = (M, num_k_blocks)
_quant_shuffle_kernel[grid](
A, fp4_buf, scale_buf,
A.stride(0), A.stride(1),
K_HALF=K_half, scaleN_pad=scale_n_pad,
M_VAL=M, BLOCK_K=BLOCK_K,
num_warps=4, num_stages=1,
)
A_q = fp4_buf.view(torch.float4_e2m1fn_x2)
A_scale = scale_buf.view(torch.float8_e8m0fnu)
aiter.gemm_a4w4_asm(A_q, B_shuffle, A_scale, B_scale_sh,
out, _ASM_KERNEL, bpreshuffle=True)
return out
scrolls · 391 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