submission 743751
Andrewxu313 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 358 lines, June 9 Researcher Reciprocity License v1.0.
v140h_shuffled_gemm.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-743751?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:a3a6c0cfdf0877ac4999e502ae6ed0782cde343d7e5fc1e7489ac3f300654835
license declaredunknown
license concludedunknown
authorsAndrewxu313
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
tile-m = 16
BLOCK_M=16, BLOCK_N=32,tile-n = 32
BLOCK_M=16, BLOCK_N=32,Kernel source
v140h_shuffled_gemm.py358 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""v140h: Custom Triton GEMM reading shuffled B data directly.
Uses exact 6D permutation formulas from aiter source code:
- shuffle_weight: view(1, N//16, 16, Kh//32, 2, 16).permute(0,1,3,4,2,5)
- e8m0_shuffle: view(sm//32, 2, 16, sn//8, 2, 4).permute(0,3,5,2,4,1)
with padding: sm = ceil(N/256)*256, sn = ceil(Ks/8)*8
NO dynamic_mxfp4_quant(B). NO e8m0_shuffle call. NO ASM GEMM.
"""
import torch
import triton
import triton.language as tl
# ─── FP4 E2M1 conversion ───
@triton.jit
def _to_e2m1(scaled_val):
bits = scaled_val.to(tl.uint32, bitcast=True)
sign = (bits >> 28) & 8
pb = bits & 0x7FFFFFFF
pf = pb.to(tl.float32, bitcast=True)
dm = tl.full(pb.shape, 0x4A800000, dtype=tl.uint32)
tmp_sub = pf + dm.to(tl.float32, bitcast=True)
tb = tmp_sub.to(tl.uint32, bitcast=True)
r_sub = (tb - dm) & 0xF
mo = (pb >> 22) & 1
r_norm = ((pb + 0xC11FFFFF + mo) >> 22) & 7
r = tl.where(pf < 1.0, r_sub, r_norm)
r = tl.where(pf >= 6.0, tl.full(r.shape, 7, dtype=tl.uint32), r)
return (r | sign).to(tl.uint8)
# ─── A quant kernel (raw output) ───
@triton.jit
def _quant_kernel(A_ptr, Aq_ptr, As_ptr, M, K, Kh, Ks, BLOCK: tl.constexpr):
pid = tl.program_id(0)
row = pid // Ks
bl = pid - row * Ks
if row >= M:
return
k_base = bl * 32
pair_off = tl.arange(0, BLOCK)
base = A_ptr + row * K + k_base
e0 = tl.load(base + pair_off * 2, mask=(k_base + pair_off * 2) < K, other=0.0).to(tl.float32)
e1 = tl.load(base + pair_off * 2 + 1, mask=(k_base + pair_off * 2 + 1) < K, other=0.0).to(tl.float32)
amax = tl.maximum(tl.max(tl.abs(e0), axis=0), tl.max(tl.abs(e1), axis=0))
tmp = tl.where(pair_off == 0, amax, 0.0)
amax_bits = tl.sum(tmp.to(tl.uint32, bitcast=True), axis=0)
rounded = (amax_bits + 0x200000) & 0xFF800000
biased_exp = (rounded >> 23) & 0xFF
su = tl.where(biased_exp > 0, biased_exp.to(tl.int32) - 129, tl.full([], -127, dtype=tl.int32))
su = tl.maximum(su, -127)
su = tl.minimum(su, 127)
e8 = (su + 127).to(tl.uint8)
qs = tl.math.exp2((tl.maximum(tl.minimum(127 - su, 254), 1) - 127).to(tl.float32))
packed = (_to_e2m1(e0 * qs).to(tl.uint8) & 0xF) | ((_to_e2m1(e1 * qs).to(tl.uint8) & 0xF) << 4)
tl.store(Aq_ptr + row * Kh + bl * 16 + pair_off, packed, mask=pair_off < 16)
tl.store(As_ptr + row * Ks + bl, e8)
# ─── Shuffled index helpers ───
# These compute flat byte offset into shuffled buffer for raw position (n, k).
@triton.jit
def _shuf_data_idx(n, kp, Kh_32):
"""shuffle_weight index: raw[n, kp] -> flat pos in shuffled buffer.
Forward: view(1, N//16, 16, Kh//32, 2, 16).permute(0,1,3,4,2,5)
Shuffled 6D = (0, n//16, kp//32, (kp//16)%2, n%16, kp%16)
Flat = n_blk16 * (Kh_32 * 2 * 16 * 16) + kp_blk32 * (2*16*16) + kp_half * (16*16) + n_inner * 16 + kp_inner
"""
n_blk16 = n // 16
n_inner = n % 16
kp_blk32 = kp // 32
kp_half = (kp // 16) % 2
kp_inner = kp % 16
return n_blk16 * (Kh_32 * 512) + kp_blk32 * 512 + kp_half * 256 + n_inner * 16 + kp_inner
@triton.jit
def _shuf_scale_idx(n, ks, sn, sm_32):
"""e8m0_shuffle index: raw[n, ks] -> flat pos in shuffled buffer.
Padding: sm = ceil(N/256)*256, sn_pad = ceil(Ks/8)*8.
Forward: view(sm//32, 2, 16, sn//8, 2, 4).permute(0,3,5,2,4,1)
6D raw decomposition of padded (n, ks):
d0 = n // 32, d1 = (n // 16) % 2, d2 = n % 16
d3 = ks // 8, d4 = (ks // 4) % 2, d5 = ks % 4
Shuffled 6D = (d0, d3, d5, d2, d4, d1)
Flat = d0*(sn_8 * 4*16*2*2) + d3*(4*16*2*2) + d5*(16*2*2) + d2*(2*2) + d4*2 + d1
where sn_8 = sn // 8
"""
d0 = n // 32
d1 = (n // 16) % 2
d2 = n % 16
sn_8 = sn // 8
d3 = ks // 8
d4 = (ks // 4) % 2
d5 = ks % 4
return d0 * (sn_8 * 256) + d3 * 256 + d5 * 64 + d2 * 4 + d4 * 2 + d1
# ─── GEMM with shuffled B (no-mask, pipelined) ───
@triton.jit
def _gemm_shuf_nomask(
A_ptr, B_shuf_ptr, C_ptr, A_scale_ptr, B_scale_shuf_ptr,
stride_am, stride_ak, stride_as_m, stride_as_k,
stride_cm, stride_cn,
Kh_32, Ks_sn, Ks_sm_32,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr, NUM_STAGES: tl.constexpr,
M_CONST: tl.constexpr, N_CONST: tl.constexpr, K_CONST: tl.constexpr,
):
pid = tl.program_id(0)
NUM_PID_M: tl.constexpr = M_CONST // BLOCK_M
NUM_PID_N: tl.constexpr = N_CONST // BLOCK_N
num_pid_in_group: tl.constexpr = 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)
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
BKP: tl.constexpr = BLOCK_K // 2
BKS: tl.constexpr = BLOCK_K // 32
offs_kp = tl.arange(0, BKP)
offs_ks = tl.arange(0, BKS)
NUM_K: tl.constexpr = K_CONST // BLOCK_K
for k_iter in tl.range(0, NUM_K, num_stages=NUM_STAGES):
k_start = k_iter * BLOCK_K
kp = k_start // 2
ks = k_start // 32
# A loads (raw layout)
a = tl.load(A_ptr + offs_m[:, None] * stride_am + (kp + offs_kp[None, :]) * stride_ak)
a_scale = tl.load(A_scale_ptr + offs_m[:, None] * stride_as_m + (ks + offs_ks[None, :]) * stride_as_k)
# B data from shuffled layout: b[BKP, BN]
kp_abs = kp + offs_kp # [BKP]
b_idx = _shuf_data_idx(offs_n[None, :], kp_abs[:, None], Kh_32)
b = tl.load(B_shuf_ptr + b_idx)
# B scale from shuffled layout: b_scale[BN, BKS]
ks_abs = ks + offs_ks # [BKS]
bs_idx = _shuf_scale_idx(offs_n[:, None], ks_abs[None, :], Ks_sn, Ks_sm_32)
b_scale = tl.load(B_scale_shuf_ptr + bs_idx)
acc = tl.dot_scaled(a, a_scale, "e2m1", b, b_scale, "e2m1", acc=acc)
c_ptrs = C_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
tl.store(c_ptrs, acc.to(tl.bfloat16))
# ─── GEMM with shuffled B (masked) ───
@triton.jit
def _gemm_shuf_masked(
A_ptr, B_shuf_ptr, C_ptr, A_scale_ptr, B_scale_shuf_ptr,
M, N, K,
stride_am, stride_ak, stride_as_m, stride_as_k,
stride_cm, stride_cn,
Kh_32, Ks_sn, Ks_sm_32,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr,
):
pid = tl.program_id(0)
num_pid_m = tl.cdiv(M, 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)
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
BKP: tl.constexpr = BLOCK_K // 2
BKS: tl.constexpr = BLOCK_K // 32
offs_kp = tl.arange(0, BKP)
offs_ks = tl.arange(0, BKS)
for k_start in range(0, K, BLOCK_K):
kp = k_start // 2
ks = k_start // 32
a = tl.load(A_ptr + offs_m[:, None] * stride_am + (kp + offs_kp[None, :]) * stride_ak,
mask=(offs_m[:, None] < M) & ((kp + offs_kp[None, :]) < K // 2), other=0)
a_scale = tl.load(A_scale_ptr + offs_m[:, None] * stride_as_m + (ks + offs_ks[None, :]) * stride_as_k,
mask=(offs_m[:, None] < M) & ((ks + offs_ks[None, :]) < K // 32), other=127)
kp_abs = kp + offs_kp
b_idx = _shuf_data_idx(offs_n[None, :], kp_abs[:, None], Kh_32)
b = tl.load(B_shuf_ptr + b_idx, mask=((kp_abs[:, None]) < K // 2) & (offs_n[None, :] < N), other=0)
ks_abs = ks + offs_ks
bs_idx = _shuf_scale_idx(offs_n[:, None], ks_abs[None, :], Ks_sn, Ks_sm_32)
b_scale = tl.load(B_scale_shuf_ptr + bs_idx, mask=(offs_n[:, None] < N) & ((ks_abs[None, :]) < K // 32), other=127)
acc = tl.dot_scaled(a, a_scale, "e2m1", b, b_scale, "e2m1", acc=acc)
c_ptrs = C_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
tl.store(c_ptrs, acc.to(tl.bfloat16), mask=(offs_m[:, None] < M) & (offs_n[None, :] < N))
# ─── Fused K=512 with shuffled B ───
@triton.jit
def _fused_k512_shuf(
A_bf16_ptr, B_shuf_ptr, C_ptr, B_scale_shuf_ptr,
M, N, Kh_32, Ks_sn, Ks_sm_32,
stride_am, stride_ak, stride_cm, stride_cn,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
):
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)
BK_PACKED: tl.constexpr = 256
BK_SCALES: tl.constexpr = 16
PAIRS_PER_GROUP: tl.constexpr = 16
offs_kp = tl.arange(0, BK_PACKED)
mask_m = offs_m[:, None] < M
a_e_ptrs = A_bf16_ptr + offs_m[:, None] * stride_am + (offs_kp[None, :] * 2) * stride_ak
a_o_ptrs = A_bf16_ptr + offs_m[:, None] * stride_am + (offs_kp[None, :] * 2 + 1) * stride_ak
a_even = tl.load(a_e_ptrs, mask=mask_m, other=0.0).to(tl.float32)
a_odd = tl.load(a_o_ptrs, mask=mask_m, other=0.0).to(tl.float32)
abs_all = tl.maximum(tl.abs(a_even), tl.abs(a_odd))
grouped = tl.reshape(abs_all, (BLOCK_M * BK_SCALES, PAIRS_PER_GROUP))
amax = tl.max(grouped, axis=1)
amax_bits = amax.to(tl.uint32, bitcast=True)
rounded = (amax_bits + 0x200000) & 0xFF800000
biased_exp = (rounded >> 23) & 0xFF
su = tl.where(biased_exp > 0, biased_exp.to(tl.int32) - 129, -127)
su = tl.maximum(su, -127)
su = tl.minimum(su, 127)
e8 = (su + 127).to(tl.uint8)
a_scale = tl.reshape(e8, (BLOCK_M, BK_SCALES))
qs_exp = tl.maximum(tl.minimum(127 - su, 254), 1)
qs = tl.math.exp2((qs_exp - 127).to(tl.float32))
qs_2d = tl.reshape(qs, (BLOCK_M * BK_SCALES, 1))
qs_bc = tl.broadcast_to(qs_2d, (BLOCK_M * BK_SCALES, PAIRS_PER_GROUP))
qs_flat = tl.reshape(qs_bc, (BLOCK_M, BK_PACKED))
fp4_e = _to_e2m1(a_even * qs_flat)
fp4_o = _to_e2m1(a_odd * qs_flat)
a_q = (fp4_e.to(tl.uint8) & 0xF) | ((fp4_o.to(tl.uint8) & 0xF) << 4)
# B from shuffled
offs_ks = tl.arange(0, BK_SCALES)
b_idx = _shuf_data_idx(offs_n[None, :], offs_kp[:, None], Kh_32)
b = tl.load(B_shuf_ptr + b_idx, mask=offs_n[None, :] < N, other=0)
bs_idx = _shuf_scale_idx(offs_n[:, None], offs_ks[None, :], Ks_sn, Ks_sm_32)
b_scale = tl.load(B_scale_shuf_ptr + bs_idx, mask=offs_n[:, None] < N, other=127)
acc = tl.dot_scaled(a_q, a_scale, "e2m1", b, b_scale, "e2m1")
c_ptrs = C_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
tl.store(c_ptrs, acc.to(tl.bfloat16), mask=(offs_m[:, None] < M) & (offs_n[None, :] < N))
# ─── Dispatch ───
_cache = {}
def custom_kernel(data):
A = data[0]
B_shuffle = data[3]
B_scale_sh = data[4]
m, k = A.shape
n = B_shuffle.shape[0]
Kh = k // 2
Ks = k // 32
# Shuffle params
Kh_32 = Kh // 32 # for shuffle_weight: number of 32-byte blocks in Kh
# e8m0_shuffle padding
sm_pad = ((n + 255) // 256) * 256
sn_pad = ((Ks + 7) // 8) * 8
Ks_sm_32 = sm_pad // 32 # for scale shuffle formula
B_sh_flat = B_shuffle.view(torch.uint8).reshape(-1)
Bs_sh_flat = B_scale_sh.view(torch.uint8).reshape(-1)
key = (m, n, k)
if key not in _cache:
dev = A.device
_cache[key] = (
torch.zeros(m, Kh, dtype=torch.uint8, device=dev),
torch.zeros(m, Ks, dtype=torch.uint8, device=dev),
torch.empty(m, n, dtype=torch.bfloat16, device=dev),
)
Aq, As, C = _cache[key]
if k == 512 and m <= 32:
grid = ((m + 15) // 16, (n + 31) // 32)
_fused_k512_shuf[grid](
A, B_sh_flat, C, Bs_sh_flat,
m, n, Kh_32, sn_pad, Ks_sm_32,
A.stride(0), A.stride(1), C.stride(0), C.stride(1),
BLOCK_M=16, BLOCK_N=32,
)
return C
_quant_kernel[(m * Ks,)](A, Aq, As, m, k, Kh, Ks, BLOCK=16)
# Shape dispatch
if k == 7168 and m == 16 and n == 2112:
BM, BN, BK, GSM, NS = 16, 16, 1024, 1, 4
elif k == 7168:
BM, BN, BK, GSM, NS = 16, 16, 1024, 1, 3
elif k == 2048 and m == 64 and n == 7168:
BM, BN, BK, GSM, NS = 64, 32, 512, 4, 3
elif k == 2048:
BM, BN, BK, GSM, NS = 64, 32, 512, 4, 3
elif k == 1536 and m == 256 and n == 3072:
BM, BN, BK, GSM, NS = 128, 32, 512, 2, 3
elif k == 512:
BM, BN, BK, GSM, NS = 32, 64, 64, 4, 1
else:
BM, BN, BK, GSM, NS = 64, 64, 64, 8, 1
if m % BM == 0 and n % BN == 0:
grid = ((m // BM) * (n // BN),)
_gemm_shuf_nomask[grid](
Aq, B_sh_flat, C, As, Bs_sh_flat,
Aq.stride(0), Aq.stride(1), As.stride(0), As.stride(1),
C.stride(0), C.stride(1),
Kh_32, sn_pad, Ks_sm_32,
BLOCK_M=BM, BLOCK_N=BN, BLOCK_K=BK,
GROUP_SIZE_M=GSM, NUM_STAGES=NS,
M_CONST=m, N_CONST=n, K_CONST=k,
)
else:
grid = (triton.cdiv(m, BM) * triton.cdiv(n, BN),)
_gemm_shuf_masked[grid](
Aq, B_sh_flat, C, As, Bs_sh_flat,
m, n, k,
Aq.stride(0), Aq.stride(1), As.stride(0), As.stride(1),
C.stride(0), C.stride(1),
Kh_32, sn_pad, Ks_sm_32,
BLOCK_M=BM, BLOCK_N=BN, BLOCK_K=BK,
GROUP_SIZE_M=GSM,
)
return C
scrolls · 358 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