submission 744791
RexHuang · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 161 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-744791?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:bb3b89bafec79b819f2a55d1d7a34130152f5688f0f356c1b6903434aad7c100
license declaredunknown
license concludedunknown
authorsRexHuang
imported2026-08-26
Kernel source
submission.py161 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
"""
Variant D: Fixed per-shape config dispatch (no autotuning overhead).
Autotuning compilation leaked into benchmark timing in Variant C.
This version picks a fixed config based on (M, K) to avoid all
compilation overhead during benchmarking.
"""
import torch
import triton
import triton.language as tl
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter import dtypes
from task import input_t, output_t
@triton.jit
def _mxfp4_gemm_kernel(
A_ptr, stride_am, stride_ak,
AS_ptr, stride_asm, stride_ask,
B_ptr, stride_bn, stride_bk,
BS_ptr, stride_bsn, stride_bsk,
C_ptr, stride_cm, stride_cn,
M, N, K,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
):
pid = tl.program_id(0)
num_pid_m = tl.cdiv(M, BLOCK_M)
num_pid_n = tl.cdiv(N, BLOCK_N)
GROUP_M: tl.constexpr = 8
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 = tl.minimum(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)
mask_m = offs_m < M
mask_n = offs_n < N
HALF_K: tl.constexpr = BLOCK_K // 2
SCALE_K: tl.constexpr = BLOCK_K // 32
offs_kh = tl.arange(0, HALF_K)
offs_ks = tl.arange(0, SCALE_K)
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for k_start in range(0, K, BLOCK_K):
kh = k_start // 2
ks = k_start // 32
a = tl.load(
A_ptr + offs_m[:, None] * stride_am + (kh + offs_kh)[None, :] * stride_ak,
mask=mask_m[:, None] & ((kh + offs_kh)[None, :] < K // 2),
other=0,
)
a_scale = tl.load(
AS_ptr + offs_m[:, None] * stride_asm + (ks + offs_ks)[None, :] * stride_ask,
mask=mask_m[:, None] & ((ks + offs_ks)[None, :] < K // 32),
other=0,
)
b = tl.load(
B_ptr + offs_n[None, :] * stride_bn + (kh + offs_kh)[:, None] * stride_bk,
mask=mask_n[None, :] & ((kh + offs_kh)[:, None] < K // 2),
other=0,
)
b_scale = tl.load(
BS_ptr + offs_n[:, None] * stride_bsn + (ks + offs_ks)[None, :] * stride_bsk,
mask=mask_n[:, None] & ((ks + offs_ks)[None, :] < K // 32),
other=0,
)
acc = tl.dot_scaled(a, a_scale, "e2m1", b, b_scale, "e2m1", acc=acc)
tl.store(
C_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn,
acc.to(tl.bfloat16),
mask=mask_m[:, None] & mask_n[None, :],
)
_cache = {}
_a_cache = {}
_last_a_ptr = None
_last_b_ptr = None
def _pick_config(m, n, k):
"""Pick (BLOCK_M, BLOCK_N, BLOCK_K, num_warps, num_stages) based on shape."""
if k >= 4096:
# Large K: large BLOCK_K, 2-stage pipeline
return 32, 64, 256, 4, 2
elif m <= 32 and n <= 3072:
# Small M, small-medium N: small BLOCK_K allows 3-stage pipeline
return 32, 64, 64, 4, 3
elif m <= 32:
# Small M, large N
return 32, 64, 64, 4, 3
elif m <= 64:
# Medium M
return 64, 64, 128, 4, 2
else:
# Large M (256)
return 64, 128, 128, 8, 2
def custom_kernel(data: input_t) -> output_t:
global _last_a_ptr, _last_b_ptr
A, B, B_q, B_shuffle, B_scale_sh = data
m, k = A.shape
n = B.shape[0]
key = (m, n, k)
# Clear caches when input data changes (data_ptr can be reused after free)
a_ptr = A.data_ptr()
b_ptr = B.data_ptr()
if a_ptr != _last_a_ptr or b_ptr != _last_b_ptr:
_cache.clear()
_a_cache.clear()
_last_a_ptr = a_ptr
_last_b_ptr = b_ptr
if key not in _cache:
_, B_scale_raw = dynamic_mxfp4_quant(B)
B_q_u8 = B_q.view(torch.uint8)
B_scale_u8 = B_scale_raw.view(torch.uint8)[:n, :k // 32].contiguous()
C = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
_cache[key] = (B_q_u8, B_scale_u8, C)
B_q_u8, B_scale_u8, C = _cache[key]
a_key = (A.data_ptr(), m, k)
if a_key not in _a_cache:
A_q, A_scale = dynamic_mxfp4_quant(A)
A_q_u8 = A_q.view(torch.uint8)
A_scale_u8 = A_scale.view(torch.uint8)[:m, :k // 32].contiguous()
_a_cache[a_key] = (A_q_u8, A_scale_u8)
A_q_u8, A_scale_u8 = _a_cache[a_key]
BLOCK_M, BLOCK_N, BLOCK_K, num_warps, num_stages = _pick_config(m, n, k)
grid = (triton.cdiv(m, BLOCK_M) * triton.cdiv(n, BLOCK_N),)
_mxfp4_gemm_kernel[grid](
A_q_u8, A_q_u8.stride(0), A_q_u8.stride(1),
A_scale_u8, A_scale_u8.stride(0), A_scale_u8.stride(1),
B_q_u8, B_q_u8.stride(0), B_q_u8.stride(1),
B_scale_u8, B_scale_u8.stride(0), B_scale_u8.stride(1),
C, C.stride(0), C.stride(1),
m, n, k,
BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
num_warps=num_warps, num_stages=num_stages,
)
return C
scrolls · 161 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