submission 736751
sepehresy · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 238 lines, June 9 Researcher Reciprocity License v1.0.
MM_V68.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-736751?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:ec79f54f24afb15a856f6f80d34f85b08d2ccd20b8b8931f526e2d02242b1d84
license declaredunknown
license concludedunknown
authorssepehresy
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
Kernel source
MM_V68.py238 lines
"""
MM_V68: V67 + deeper pipeline for the main K=1536 M=256 bottleneck.
V67 fixed KernelGuard issues and improved the K=512/K=7168 rows, but the
ranked leaderboard still shows the slow tail centered on:
(256, 3072, 1536) ~17.8us
This version keeps V67's KernelGuard-safe structure and only bumps
num_stages from 2 -> 3 for that one shape to test whether the 3 x BLOCK_K
pipeline benefits from deeper software pipelining without changing the
winning configs elsewhere.
"""
import os
os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
os.environ["GPU_MAX_HW_QUEUES"] = "2"
import torch
import triton
import triton.language as tl
from task import input_t, output_t
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import _mxfp4_quant_op
_u8 = torch.uint8
_bf16 = torch.bfloat16
_f32 = torch.float32
@triton.jit
def _fused_gemm_kernel(
a_ptr, b_ptr, c_ptr, bs_ptr,
M, N, K,
stride_am, stride_ak,
stride_bk, stride_bn,
stride_ck, stride_cm, stride_cn,
SN,
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr, NUM_KSPLIT: tl.constexpr,
K_ITERS_PER_SPLIT: tl.constexpr,
EVEN_MNK: tl.constexpr,
num_warps: tl.constexpr, num_stages: tl.constexpr,
waves_per_eu: tl.constexpr,
):
K_HALF: tl.constexpr = BLOCK_K // 2
SCALE_K: tl.constexpr = BLOCK_K // 32
SN32 = SN * 32
GRID_MN = tl.cdiv(M, BLOCK_M) * tl.cdiv(N, BLOCK_N)
pid_unified = tl.program_id(0)
pid_k = pid_unified // GRID_MN
pid = pid_unified % GRID_MN
num_pid_m = tl.cdiv(M, BLOCK_M)
num_pid_n = tl.cdiv(N, BLOCK_N)
if NUM_KSPLIT == 1:
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 = tl.minimum(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
else:
pid_m = pid // num_pid_n
pid_n = pid % 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)
sn_i0 = offs_n // 32
sn_i1 = (offs_n >> 4) & 1
sn_i2 = offs_n & 15
src_r = sn_i0 * SN32 + sn_i1 + sn_i2 * 4
a_base = a_ptr + offs_m[:, None] * stride_am
accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
k_start_split = pid_k * K_ITERS_PER_SPLIT * BLOCK_K
for ki in range(K_ITERS_PER_SPLIT):
k_start = k_start_split + ki * BLOCK_K
if k_start < K:
k_half_start = k_start // 2
a_k_offs = k_start + tl.arange(0, BLOCK_K)
b_k_offs = k_half_start + tl.arange(0, BLOCK_K // 2)
if EVEN_MNK:
a_tile = tl.load(a_base + a_k_offs[None, :] * stride_ak)
b_tile = tl.load(b_ptr + b_k_offs[:, None] * stride_bk + offs_n[None, :] * stride_bn)
else:
a_mask = (offs_m[:, None] < M) & (a_k_offs[None, :] < K)
a_tile = tl.load(a_base + a_k_offs[None, :] * stride_ak, mask=a_mask, other=0.0)
b_mask = (b_k_offs[:, None] < (K // 2)) & (offs_n[None, :] < N)
b_tile = tl.load(
b_ptr + b_k_offs[:, None] * stride_bk + offs_n[None, :] * stride_bn,
mask=b_mask,
other=0,
)
scale_k_idx = k_start // 32
sk_offs = scale_k_idx + tl.arange(0, SCALE_K)
sk_i3 = sk_offs // 8
sk_i4 = (sk_offs >> 2) & 1
sk_i5 = sk_offs & 3
src_c = sk_i3 * 256 + sk_i4 * 2 + sk_i5 * 64
bs_src = src_r[:, None] + src_c[None, :]
if EVEN_MNK:
b_scales = tl.load(bs_ptr + bs_src)
else:
bs_mask = (offs_n[:, None] < N) & (sk_offs[None, :] < (K // 32))
b_scales = tl.load(bs_ptr + bs_src, mask=bs_mask, other=0)
a_fp4, a_scales = _mxfp4_quant_op(a_tile, BLOCK_K, BLOCK_M, 32)
accumulator = tl.dot_scaled(a_fp4, a_scales, "e2m1", b_tile, b_scales, "e2m1", accumulator)
c = accumulator.to(tl.bfloat16) if NUM_KSPLIT == 1 else accumulator
c_mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
c_ptrs = c_ptr + pid_k * stride_ck + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
tl.store(c_ptrs, c, mask=c_mask)
@triton.jit
def _reduce_kernel(
partials_ptr, out_ptr, M, N,
stride_pk, stride_pm, stride_pn,
stride_om, stride_on,
NUM_SPLITS: tl.constexpr,
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)
mask = (offs_m[:, None] < M) & (offs_n[None, :] < N)
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
for k in range(NUM_SPLITS):
p = tl.load(
partials_ptr + k * stride_pk + offs_m[:, None] * stride_pm + offs_n[None, :] * stride_pn,
mask=mask,
other=0.0,
)
acc += p
tl.store(
out_ptr + offs_m[:, None] * stride_om + offs_n[None, :] * stride_on,
acc.to(tl.bfloat16),
mask=mask,
)
# (BLOCK_M, NUM_KSPLIT, K_ITERS, num_warps, num_stages, EVEN_MNK)
_SHAPE_CFGS = {
(4, 2880, 512): (16, 1, 1, 8, 2, False),
(32, 4096, 512): (16, 1, 1, 8, 2, True),
(32, 2880, 512): (16, 1, 1, 8, 2, False),
(256, 2880, 512): (16, 1, 1, 8, 2, False),
(8, 2112, 7168): (16, 14, 1, 4, 1, False),
(16, 2112, 7168): (16, 14, 1, 4, 1, False),
(64, 7168, 2048): (16, 2, 2, 4, 2, True),
(256, 3072, 1536): (16, 1, 3, 4, 3, True),
(16, 3072, 1536): (16, 1, 3, 4, 2, True),
(64, 3072, 1536): (16, 1, 3, 4, 2, True),
}
def custom_kernel(data: input_t) -> output_t:
A = data[0]
B_q = data[2]
B_scale_sh = data[4]
if not A.is_contiguous():
A = A.contiguous()
m, k = A.shape
n = B_q.shape[0]
B_u8 = B_q.view(_u8)
bs_u8 = B_scale_sh.view(_u8)
sn = bs_u8.shape[1]
BLOCK_N = 128
BLOCK_K = 512
cfg = _SHAPE_CFGS.get((m, n, k))
if cfg is not None:
BLOCK_M, NUM_KSPLIT, K_ITERS, nw, ns, even = cfg
else:
BLOCK_M = 16
NUM_KSPLIT = max(1, k // BLOCK_K) if k >= 4096 else 1
K_ITERS = max(1, triton.cdiv(k, max(NUM_KSPLIT, 1) * BLOCK_K))
nw, ns, even = 4, 1, False
grid_mn = triton.cdiv(m, BLOCK_M) * triton.cdiv(n, BLOCK_N)
if NUM_KSPLIT <= 1:
out = torch.empty((m, n), dtype=_bf16, device=A.device)
_fused_gemm_kernel[(grid_mn,)](
A, B_u8, out, bs_u8,
m, n, k,
A.stride(0), A.stride(1),
B_u8.stride(1), B_u8.stride(0),
0, out.stride(0), out.stride(1),
sn,
BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
GROUP_SIZE_M=4, NUM_KSPLIT=1,
K_ITERS_PER_SPLIT=K_ITERS,
EVEN_MNK=even,
num_warps=nw, num_stages=ns, waves_per_eu=2,
)
return out
partials = torch.empty((NUM_KSPLIT, m, n), dtype=_f32, device=A.device)
out = torch.empty((m, n), dtype=_bf16, device=A.device)
_fused_gemm_kernel[(NUM_KSPLIT * grid_mn,)](
A, B_u8, partials, bs_u8,
m, n, k,
A.stride(0), A.stride(1),
B_u8.stride(1), B_u8.stride(0),
partials.stride(0), partials.stride(1), partials.stride(2),
sn,
BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K,
GROUP_SIZE_M=4, NUM_KSPLIT=NUM_KSPLIT,
K_ITERS_PER_SPLIT=K_ITERS,
EVEN_MNK=even,
num_warps=nw, num_stages=ns, waves_per_eu=2,
)
_reduce_kernel[(triton.cdiv(m, 16), triton.cdiv(n, 64))](
partials, out, m, n,
partials.stride(0), partials.stride(1), partials.stride(2),
out.stride(0), out.stride(1),
NUM_SPLITS=NUM_KSPLIT,
BLOCK_M=16, BLOCK_N=64,
)
return out
scrolls · 238 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