submission 577240
Arseni Ivanov · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 315 lines, June 9 Researcher Reciprocity License v1.0.
submission_7.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-577240?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:3463a073ca833e7abe0ebd255c52a457572fb66e99d4a11dbbf9e5e6a6316729
license declaredunknown
license concludedunknown
authorsArseni Ivanov
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
submission_7.py315 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
@triton.jit
def inline_quantize_mxfp4_optimized(x, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_K: tl.constexpr):
"""Optimized MXFP4 quantization with reduced operations"""
x_reshaped = tl.reshape(x, [BLOCK_SIZE_M, BLOCK_SIZE_K // 32, 32])
amax = tl.max(tl.abs(x_reshaped), axis=2)
amax_bits = amax.to(tl.uint32, bitcast=True)
amax_rounded = ((amax_bits + 0x200000) & 0xFF800000)
exp_biased = (amax_rounded >> 23).to(tl.int32)
scale_unbiased = tl.minimum(tl.maximum(exp_biased - 129, -127), 127)
quant_scale_bits = ((127 - scale_unbiased).to(tl.uint32) << 23)
quant_scale = quant_scale_bits.to(tl.float32, bitcast=True)
qx = x_reshaped * tl.reshape(quant_scale, [BLOCK_SIZE_M, BLOCK_SIZE_K // 32, 1])
qx_u32 = qx.to(tl.uint32, bitcast=True)
sign = qx_u32 & 0x80000000
qx_abs = (qx_u32 ^ sign).to(tl.float32, bitcast=True)
denorm_val = ((qx_abs + 4194304.0).to(tl.int32, bitcast=True) - 0x4A800000)
norm_bits = qx_abs.to(tl.int32, bitcast=True)
mant_lsb = (norm_bits >> 22) & 1
norm_val = ((norm_bits + (-1054867457 + mant_lsb)) >> 22)
e2m1 = tl.where(qx_abs < 1.0, denorm_val, norm_val)
e2m1 = tl.where(qx_abs >= 6.0, 7, e2m1)
e2m1_packed = ((sign >> 28).to(tl.int32) | e2m1).to(tl.uint8)
e2m1_pairs = tl.reshape(e2m1_packed, [BLOCK_SIZE_M, BLOCK_SIZE_K // 2, 2])
evens, odds = tl.split(e2m1_pairs)
x_fp4 = evens | (odds << 4)
bs_e8m0 = (scale_unbiased + 127).to(tl.uint8)
return x_fp4, bs_e8m0
HARDCODED_CONFIGS = {
(4, 2880, 512): {'BLOCK_M': 16, 'BLOCK_N': 64, 'BLOCK_K': 512, 'GROUP_M': 8, 'LOOP_STAGES': 2, 'num_warps': 4, 'num_stages': 1},
(16, 2112, 7168): {'BLOCK_M': 16, 'BLOCK_N': 128, 'BLOCK_K': 256, 'GROUP_M': 8, 'LOOP_STAGES': 2, 'num_warps': 4, 'num_stages': 2},
(32, 4096, 512): {'BLOCK_M': 16, 'BLOCK_N': 32, 'BLOCK_K': 256, 'GROUP_M': 8, 'LOOP_STAGES': 2, 'num_warps': 4, 'num_stages': 2},
(32, 2880, 512): {'BLOCK_M': 16, 'BLOCK_N': 64, 'BLOCK_K': 512, 'GROUP_M': 8, 'LOOP_STAGES': 2, 'num_warps': 4, 'num_stages': 1},
(64, 7168, 2048): {'BLOCK_M': 16, 'BLOCK_N': 256, 'BLOCK_K': 256, 'GROUP_M': 8, 'LOOP_STAGES': 2, 'num_warps': 8, 'num_stages': 2},
(256, 3072, 1536): {'BLOCK_M': 16, 'BLOCK_N': 256, 'BLOCK_K': 256, 'GROUP_M': 8, 'LOOP_STAGES': 2, 'num_warps': 8, 'num_stages': 2},
}
def get_kernel_config(m, n, k):
if (m, n, k) in HARDCODED_CONFIGS:
return HARDCODED_CONFIGS[(m, n, k)]
if k >= 4096:
return {'BLOCK_M': 16, 'BLOCK_N': 128, 'BLOCK_K': 256, 'GROUP_M': 8, 'LOOP_STAGES': 2, 'num_warps': 4, 'num_stages': 2}
elif k <= 512:
return {'BLOCK_M': 16, 'BLOCK_N': 64, 'BLOCK_K': 512, 'GROUP_M': 8, 'LOOP_STAGES': 2, 'num_warps': 4, 'num_stages': 1}
else:
return {'BLOCK_M': 16, 'BLOCK_N': 128, 'BLOCK_K': 256, 'GROUP_M': 8, 'LOOP_STAGES': 2, 'num_warps': 8, 'num_stages': 2}
HARDCODED_REDUCE_CONFIGS = {
(16, 2112): {'BLOCK_M': 16, 'BLOCK_N': 64, 'num_warps': 4},
(64, 7168): {'BLOCK_M': 32, 'BLOCK_N': 64, 'num_warps': 8},
}
def get_reduce_config(m, n):
if (m, n) in HARDCODED_REDUCE_CONFIGS:
return HARDCODED_REDUCE_CONFIGS[(m, n)]
if m >= 64:
return {'BLOCK_M': 32, 'BLOCK_N': 64, 'num_warps': 8}
else:
return {'BLOCK_M': 16, 'BLOCK_N': 64, 'num_warps': 4}
# Removed @triton.heuristics completely!
@triton.jit
def fused_mxfp4_dot_scaled_kernel(
A_ptr, B_ptr, B_scale_ptr, C_ptr, Workspace_ptr,
M, N, K,
stride_am, stride_ak,
stride_bn, stride_bk,
stride_bsn, stride_bsk,
stride_cm, stride_cn,
SPLIT_K: tl.constexpr,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, GROUP_M: tl.constexpr,
EVEN_M: tl.constexpr, EVEN_N: tl.constexpr, EVEN_K: tl.constexpr, LOOP_STAGES: tl.constexpr,
USE_SHUFFLED_B: tl.constexpr, EVICT_B_FIRST: tl.constexpr, USE_CG_C: tl.constexpr
):
pid = tl.program_id(axis=0)
pid_k = tl.program_id(axis=1)
num_pid_m = tl.cdiv(M, BLOCK_M)
num_pid_n = tl.cdiv(N, BLOCK_N)
if K == 512:
pid_m = pid % num_pid_m
pid_n = pid // num_pid_m
else:
num_pid_in_group = GROUP_M * num_pid_n
group_id = pid // num_pid_in_group
first_pid_m = group_id * GROUP_M
group_size_m = min(num_pid_m - first_pid_m, GROUP_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)
a_ptrs = A_ptr + (offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak)
if USE_SHUFFLED_B:
offs_bn_shuf = pid_n * (BLOCK_N // 16) + tl.arange(0, BLOCK_N // 16)
offs_k_shuf = tl.arange(0, (BLOCK_K // 2) * 16)
b_ptrs = B_ptr + (offs_bn_shuf[:, None] * stride_bn + offs_k_shuf[None, :] * stride_bk)
else:
offs_bk_q = tl.arange(0, BLOCK_K // 2)
b_ptrs = B_ptr + (offs_bk_q[:, None] * stride_bk + offs_n[None, :] * stride_bn)
offs_bsn = pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)
offs_bks = tl.arange(0, (BLOCK_K // 32) * 32)
b_scale_ptrs = B_scale_ptr + (offs_bsn[:, None] * stride_bsn + offs_bks[None, :] * stride_bsk)
mask_m = offs_m < M
mask_n = offs_n < N
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
a_ptrs += pid_k * BLOCK_K * stride_ak
b_ptrs += pid_k * ((BLOCK_K // 2) * (16 if USE_SHUFFLED_B else 1)) * stride_bk
b_scale_ptrs += pid_k * BLOCK_K * stride_bsk
total_k_blocks = tl.cdiv(K, BLOCK_K)
for k_idx in tl.range(pid_k, total_k_blocks, SPLIT_K, LOOP_STAGES):
if EVEN_M and EVEN_K:
a = tl.load(a_ptrs, eviction_policy="evict_last")
elif EVEN_K:
a = tl.load(a_ptrs, mask=mask_m[:, None], other=0.0, eviction_policy="evict_last")
else:
k_mask = (k_idx * BLOCK_K + offs_k) < K
a = tl.load(a_ptrs, mask=(mask_m[:, None] & k_mask[None, :]), other=0.0, eviction_policy="evict_last")
a_q, a_scale = inline_quantize_mxfp4_optimized(a.to(tl.float32), BLOCK_M, BLOCK_K)
if USE_SHUFFLED_B:
mask_bn = (pid_n * (BLOCK_N // 16) + tl.arange(0, BLOCK_N // 16)) < (N // 16)
if EVICT_B_FIRST:
if EVEN_N and EVEN_K:
b_raw = tl.load(b_ptrs, eviction_policy="evict_first")
elif EVEN_K:
b_raw = tl.load(b_ptrs, mask=mask_bn[:, None], other=0, eviction_policy="evict_first")
else:
b_raw = tl.load(b_ptrs, mask=mask_bn[:, None], other=0, eviction_policy="evict_first")
else:
if EVEN_N and EVEN_K:
b_raw = tl.load(b_ptrs)
elif EVEN_K:
b_raw = tl.load(b_ptrs, mask=mask_bn[:, None], other=0)
else:
b_raw = tl.load(b_ptrs, mask=mask_bn[:, None], other=0)
b_q = (b_raw.reshape(1, BLOCK_N // 16, BLOCK_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5).reshape(BLOCK_N, BLOCK_K // 2).trans(1, 0))
else:
if EVICT_B_FIRST:
if EVEN_N and EVEN_K:
b_q = tl.load(b_ptrs, eviction_policy="evict_first")
elif EVEN_K:
b_q = tl.load(b_ptrs, mask=mask_n[None, :], other=0, eviction_policy="evict_first")
else:
bk_offs = (k_idx * BLOCK_K) // 2 + tl.arange(0, BLOCK_K // 2)
b_q = tl.load(b_ptrs, mask=(mask_n[None, :] & (bk_offs[:, None] < K // 2)), other=0, eviction_policy="evict_first")
else:
if EVEN_N and EVEN_K:
b_q = tl.load(b_ptrs)
elif EVEN_K:
b_q = tl.load(b_ptrs, mask=mask_n[None, :], other=0)
else:
bk_offs = (k_idx * BLOCK_K) // 2 + tl.arange(0, BLOCK_K // 2)
b_q = tl.load(b_ptrs, mask=(mask_n[None, :] & (bk_offs[:, None] < K // 2)), other=0)
mask_bsn = (offs_bsn < (N // 32))
if EVICT_B_FIRST:
if EVEN_N and EVEN_K:
b_scale_raw = tl.load(b_scale_ptrs, eviction_policy="evict_first")
else:
b_scale_raw = tl.load(b_scale_ptrs, mask=mask_bsn[:, None], other=0, eviction_policy="evict_first")
else:
if EVEN_N and EVEN_K:
b_scale_raw = tl.load(b_scale_ptrs)
else:
b_scale_raw = tl.load(b_scale_ptrs, mask=mask_bsn[:, None], other=0)
b_scale = (b_scale_raw.reshape(BLOCK_N // 32, BLOCK_K // 256, 4, 16, 2, 2, 1)
.permute(0, 5, 3, 1, 4, 2, 6).reshape(BLOCK_N, BLOCK_K // 32))
acc = tl.dot_scaled(a_q, a_scale, 'e2m1', b_q, b_scale, 'e2m1', acc)
a_ptrs += SPLIT_K * BLOCK_K * stride_ak
b_ptrs += SPLIT_K * ((BLOCK_K // 2) * (16 if USE_SHUFFLED_B else 1)) * stride_bk
b_scale_ptrs += SPLIT_K * BLOCK_K * stride_bsk
if SPLIT_K == 1:
c_ptrs = C_ptr + (offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn)
c_mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
tl.store(c_ptrs, acc.to(tl.bfloat16), mask=c_mask, cache_modifier=".cg" if USE_CG_C else "")
else:
ws_ptrs = Workspace_ptr + pid_k * (M * N) + offs_m[:, None] * N + offs_n[None, :]
ws_mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
tl.store(ws_ptrs, acc, mask=ws_mask, cache_modifier=".cg" if USE_CG_C else "")
# Removed @triton.heuristics here too
@triton.jit
def reduce_kernel(
Workspace_ptr, C_ptr, M, N, stride_cm, stride_cn,
SPLIT_K: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, USE_CG_C: tl.constexpr
):
pid_m, pid_n = tl.program_id(0), 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)
mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for k in range(SPLIT_K):
ptr = Workspace_ptr + k * M * N + offs_m[:, None] * N + offs_n[None, :]
acc += tl.load(ptr, mask=mask, other=0.0)
c_ptrs = C_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
tl.store(c_ptrs, acc.to(tl.bfloat16), mask=mask, cache_modifier=".cg" if USE_CG_C else "")
#used to avoid re-initializing memory, not used to store or return results
_workspace_cache = {}
_c_cache = {}
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
m, k = A.shape
n, _ = B.shape
use_shuffled_b = (m == 64 and n == 7168 and k == 2048) or m == 256
if use_shuffled_b:
B_in_u8 = B_shuffle.view(torch.uint8)
stride_bn = (k // 2) * 16
stride_bk = 1
else:
B_in_u8 = B_q.view(torch.uint8)
stride_bn = B_in_u8.stride(0)
stride_bk = B_in_u8.stride(1)
B_scale_sh_u8 = B_scale_sh.view(torch.uint8)
stride_bsn = k
stride_bsk = 1
split_k = 16 if k >= 7168 else (2 if k == 2048 else 1)
workspace = None
if split_k > 1:
ws_key = (A.device, m, n, split_k)
if ws_key not in _workspace_cache:
_workspace_cache[ws_key] = torch.empty((split_k, m, n), dtype=torch.float32, device=A.device)
workspace = _workspace_cache[ws_key]
c_key = (A.device, m, n)
if c_key not in _c_cache:
_c_cache[c_key] = torch.empty((m, n), device=A.device, dtype=torch.bfloat16)
C = _c_cache[c_key]
cfg = get_kernel_config(m, n, k)
grid_m = (m + cfg['BLOCK_M'] - 1) // cfg['BLOCK_M']
grid_n = (n + cfg['BLOCK_N'] - 1) // cfg['BLOCK_N']
grid = (grid_m * grid_n, split_k)
EVEN_M = (m % cfg['BLOCK_M'] == 0)
EVEN_N = (n % cfg['BLOCK_N'] == 0)
EVEN_K = (k % cfg['BLOCK_K'] == 0)
EVICT_B_FIRST = (k >= 1536)
USE_CG_C = ((m * n) > 128 * 128)
fused_mxfp4_dot_scaled_kernel[grid](
A, B_in_u8, B_scale_sh_u8, C, workspace, m, n, k,
A.stride(0), A.stride(1),
stride_bn, stride_bk,
stride_bsn, stride_bsk,
C.stride(0), C.stride(1),
SPLIT_K=split_k,
BLOCK_M=cfg['BLOCK_M'], BLOCK_N=cfg['BLOCK_N'], BLOCK_K=cfg['BLOCK_K'],
GROUP_M=cfg['GROUP_M'], LOOP_STAGES=cfg['LOOP_STAGES'],
EVEN_M=EVEN_M, EVEN_N=EVEN_N, EVEN_K=EVEN_K,
USE_SHUFFLED_B=use_shuffled_b, EVICT_B_FIRST=EVICT_B_FIRST, USE_CG_C=USE_CG_C,
num_warps=cfg['num_warps'], num_stages=cfg['num_stages']
)
if split_k > 1:
r_cfg = get_reduce_config(m, n)
r_grid_m = (m + r_cfg['BLOCK_M'] - 1) // r_cfg['BLOCK_M']
r_grid_n = (n + r_cfg['BLOCK_N'] - 1) // r_cfg['BLOCK_N']
reduce_grid = (r_grid_m, r_grid_n)
reduce_kernel[reduce_grid](
workspace, C, m, n,
C.stride(0), C.stride(1),
SPLIT_K=split_k,
BLOCK_M=r_cfg['BLOCK_M'], BLOCK_N=r_cfg['BLOCK_N'],
USE_CG_C=USE_CG_C,
num_warps=r_cfg['num_warps']
)
return C
scrolls · 315 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