submission 586303
shiyeegao · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 512 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-586303?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:005e082c8f95b30a505aa1d0efde4831ca69268757bc78b87b738080b1e642b7
license declaredunknown
license concludedunknown
authorsshiyeegao
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
num-warps = 4
num_warps=4,split-k
splitk_block_size = (k + num_ksplit - 1) // num_ksplitstages = 1
num_stages=1,Kernel source
submission.py512 lines
#!POPCORN leaderboard amd-mxfp4-mm
from __future__ import annotations
from typing import Any
_RUNTIME: dict[str, Any] | None = None
_DEFAULT_CONFIGS = (
(8, 8, 64, 256, 4, 2, 2, 1, 1, 0),
(31, 16, 64, 256, 4, 2, 2, 1, 1, 0),
(32, 32, 64, 256, 4, 2, 2, 1, 1, 0),
(64, 32, 64, 256, 4, 2, 2, 1, 1, 0),
(128, 32, 64, 256, 4, 2, 2, 1, 1, 0),
(256, 32, 64, 256, 4, 2, 2, 1, 1, 0),
(1 << 30, 32, 64, 256, 4, 2, 2, 1, 1, 0),
)
_SPECIAL_CONFIGS = {
(2112, 7168): (
(8, 8, 32, 512, 1, 2, 2, 1, 7, 2),
(31, 16, 32, 512, 1, 4, 1, 4, 14, 2),
(32, 32, 128, 512, 1, 4, 1, 1, 14, 2),
(64, 64, 32, 512, 1, 4, 2, 1, 7, 2),
(128, 64, 32, 1024, 1, 2, 2, 1, 1, 0),
(256, 64, 32, 1024, 1, 2, 2, 2, 1, 0),
(1 << 30, 256, 64, 256, 1, 2, 2, 1, 1, 0),
),
(3072, 1536): (
(8, 8, 32, 512, 1, 4, 2, 1, 1, 2),
(31, 8, 32, 512, 1, 4, 2, 1, 1, 0),
(32, 32, 32, 512, 1, 4, 2, 1, 1, 2),
(64, 64, 32, 512, 1, 2, 2, 1, 1, 2),
(128, 64, 32, 512, 1, 2, 2, 1, 1, 0),
(256, 128, 32, 512, 1, 4, 2, 1, 1, 0),
(1 << 30, 128, 64, 256, 8, 2, 2, 2, 1, 0),
),
(4096, 512): (
(8, 8, 64, 512, 1, 2, 1, 1, 1, 0),
(31, 8, 64, 512, 1, 2, 1, 1, 1, 0),
(32, 32, 64, 512, 1, 4, 1, 1, 1, 0),
(64, 32, 128, 512, 1, 4, 1, 1, 1, 0),
(128, 32, 128, 512, 1, 4, 1, 1, 1, 0),
(256, 256, 256, 512, 1, 4, 1, 1, 1, 0),
(1 << 30, 64, 256, 512, 1, 4, 1, 1, 1, 0),
),
}
def _floor_power_of_two(x: int) -> int:
if x <= 1:
return 1
return 1 << (x.bit_length() - 1)
def _pick_config(m: int, n: int, k: int) -> tuple[int, int, int, int, int, int, int, int, bool, int]:
table = _SPECIAL_CONFIGS.get((n, k), _DEFAULT_CONFIGS)
chosen = table[-1]
for entry in table:
if m <= entry[0]:
chosen = entry
break
(
_max_m,
block_m,
block_n,
block_k,
group_m,
num_warps,
num_stages,
waves_per_eu,
num_ksplit,
cache_hint,
) = chosen
splitk_block_size = (k + num_ksplit - 1) // num_ksplit
if block_k > splitk_block_size:
block_k = _floor_power_of_two(splitk_block_size)
block_k = max(block_k, 64)
even_k = (
(k % block_k == 0)
and (splitk_block_size % block_k == 0)
and (k % splitk_block_size == 0)
)
return (
block_m,
block_n,
block_k,
group_m,
num_warps,
num_stages,
waves_per_eu,
num_ksplit,
even_k,
cache_hint,
)
def _load_runtime() -> dict[str, Any]:
global _RUNTIME
if _RUNTIME is not None:
return _RUNTIME
import torch
import triton
import triton.language as tl
from aiter.ops.triton.quant import dynamic_mxfp4_quant
@triton.jit
def _remap_xcd(pid, grid_mn, NUM_XCDS: tl.constexpr = 8):
pids_per_xcd = (grid_mn + NUM_XCDS - 1) // NUM_XCDS
tall_xcds = grid_mn % NUM_XCDS
tall_xcds = tl.where(tall_xcds == 0, NUM_XCDS, tall_xcds)
xcd = pid % NUM_XCDS
local_pid = pid // NUM_XCDS
tall_pid = xcd * pids_per_xcd + local_pid
short_pid = tall_xcds * pids_per_xcd + (xcd - tall_xcds) * (pids_per_xcd - 1) + local_pid
return tl.where(xcd < tall_xcds, tall_pid, short_pid)
@triton.jit
def _pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M: tl.constexpr):
if GROUP_SIZE_M == 1:
return pid // num_pid_n, pid % num_pid_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 = tl.minimum(num_pid_m - first_pid_m, GROUP_SIZE_M)
tl.assume(group_size_m >= 0)
pid_in_group = pid % num_pid_in_group
pid_m = first_pid_m + (pid_in_group % group_size_m)
pid_n = pid_in_group // group_size_m
return pid_m, pid_n
@triton.jit
def _load_preshuffle_weight_full(ptrs, CACHE_HINT: tl.constexpr):
if CACHE_HINT == 2:
return tl.load(ptrs, cache_modifier=".cg")
if CACHE_HINT == 1:
return tl.load(ptrs, cache_modifier=".ca")
return tl.load(ptrs)
@triton.jit
def _load_preshuffle_weight_masked(ptrs, mask, CACHE_HINT: tl.constexpr):
if CACHE_HINT == 2:
return tl.load(ptrs, mask=mask, other=0, cache_modifier=".cg")
if CACHE_HINT == 1:
return tl.load(ptrs, mask=mask, other=0, cache_modifier=".ca")
return tl.load(ptrs, mask=mask, other=0)
@triton.jit
def _load_preshuffle_scale(ptrs, CACHE_HINT: tl.constexpr):
if CACHE_HINT == 2:
return tl.load(ptrs, cache_modifier=".cg")
if CACHE_HINT == 1:
return tl.load(ptrs, cache_modifier=".ca")
return tl.load(ptrs)
@triton.jit
def _mxfp4_preshuffle_kernel(
a_ptr,
b_ptr,
c_ptr,
a_scales_ptr,
b_scales_ptr,
m,
n,
k_packed,
b_rows,
a_scale_rows,
b_scale_rows,
stride_am,
stride_ak,
stride_bn,
stride_bk,
stride_ck,
stride_cm,
stride_cn,
stride_asm,
stride_ask,
stride_bsn,
stride_bsk,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
BLOCK_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr,
SMALL_M: tl.constexpr,
NUM_KSPLIT: tl.constexpr,
SPLITK_BLOCK_SIZE: tl.constexpr,
EVEN_K: tl.constexpr,
CACHE_HINT: tl.constexpr,
):
tl.assume(stride_am > 0)
tl.assume(stride_ak > 0)
tl.assume(stride_bn > 0)
tl.assume(stride_bk > 0)
tl.assume(stride_ck >= 0)
tl.assume(stride_cm > 0)
tl.assume(stride_cn > 0)
tl.assume(stride_asm > 0)
tl.assume(stride_ask > 0)
tl.assume(stride_bsn > 0)
tl.assume(stride_bsk > 0)
grid_m = tl.cdiv(m, BLOCK_M)
grid_n = tl.cdiv(n, BLOCK_N)
grid_mn = grid_m * grid_n
pid_unified = tl.program_id(axis=0)
pid_unified = _remap_xcd(pid_unified, grid_mn * NUM_KSPLIT, NUM_XCDS=8)
pid_k = pid_unified % NUM_KSPLIT
pid = pid_unified // NUM_KSPLIT
if NUM_KSPLIT == 1:
pid_m, pid_n = _pid_grid(pid, grid_m, grid_n, GROUP_SIZE_M=GROUP_SIZE_M)
else:
pid_m = pid // grid_n
pid_n = pid % grid_n
splitk_packed = SPLITK_BLOCK_SIZE // 2
k_start = pid_k * splitk_packed
if k_start < k_packed:
scale_group: tl.constexpr = 32
offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
offs_am = offs_m % m
offs_bn = (pid_n * (BLOCK_N // 16) + tl.arange(0, BLOCK_N // 16)) % b_rows
offs_bsn = (pid_n * (BLOCK_N // 32) + tl.arange(0, BLOCK_N // 32)) % b_scale_rows
acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)
num_k_iter = tl.cdiv(splitk_packed, BLOCK_K // 2)
offs_k = k_start + tl.arange(0, BLOCK_K // 2)
offs_k_shuffle = k_start * 16 + tl.arange(0, (BLOCK_K // 2) * 16)
a_ptrs = a_ptr + offs_am[:, None] * stride_am + offs_k[None, :] * stride_ak
b_ptrs = b_ptr + offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk
if SMALL_M:
a_scale_k = pid_k * (SPLITK_BLOCK_SIZE // scale_group) + tl.arange(
0, BLOCK_K // scale_group
)
a_scale_ptrs = a_scales_ptr + offs_am[:, None] * stride_asm + a_scale_k[None, :] * stride_ask
else:
offs_asm = (pid_m * (BLOCK_M // 32) + tl.arange(0, BLOCK_M // 32)) % a_scale_rows
a_scale_k = pid_k * (SPLITK_BLOCK_SIZE // scale_group) * 32 + tl.arange(
0, (BLOCK_K // scale_group) * 32
)
a_scale_ptrs = a_scales_ptr + offs_asm[:, None] * stride_asm + a_scale_k[None, :] * stride_ask
b_scale_k = pid_k * (SPLITK_BLOCK_SIZE // scale_group) * 32 + tl.arange(
0, (BLOCK_K // scale_group) * 32
)
b_scale_ptrs = b_scales_ptr + offs_bsn[:, None] * stride_bsn + b_scale_k[None, :] * stride_bsk
for k_iter in range(0, num_k_iter):
if EVEN_K:
a = tl.load(a_ptrs)
b = _load_preshuffle_weight_full(b_ptrs, CACHE_HINT=CACHE_HINT)
else:
rem_k = k_packed - (k_start + k_iter * (BLOCK_K // 2))
a = tl.load(a_ptrs, mask=tl.arange(0, BLOCK_K // 2)[None, :] < rem_k, other=0)
b = _load_preshuffle_weight_masked(
b_ptrs,
(tl.arange(0, (BLOCK_K // 2) * 16)[None, :] // 16) < rem_k,
CACHE_HINT=CACHE_HINT,
)
b = (
b.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)
)
if SMALL_M:
a_scales = tl.load(a_scale_ptrs)
else:
a_scales = (
tl.load(a_scale_ptrs)
.reshape(BLOCK_M // 32, BLOCK_K // scale_group // 8, 4, 16, 2, 2, 1)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_M, BLOCK_K // scale_group)
)
b_scales = (
_load_preshuffle_scale(b_scale_ptrs, CACHE_HINT=CACHE_HINT)
.reshape(BLOCK_N // 32, BLOCK_K // scale_group // 8, 4, 16, 2, 2, 1)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_N, BLOCK_K // scale_group)
)
acc = tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1", acc)
a_ptrs += (BLOCK_K // 2) * stride_ak
b_ptrs += (BLOCK_K // 2) * 16 * stride_bk
if SMALL_M:
a_scale_ptrs += (BLOCK_K // scale_group) * stride_ask
else:
a_scale_ptrs += BLOCK_K * stride_ask
b_scale_ptrs += BLOCK_K * stride_bsk
c = acc.to(c_ptr.type.element_ty)
c_ptrs = c_ptr + pid_k * stride_ck + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
mask = (offs_m[:, None] < m) & (offs_n[None, :] < n)
tl.store(c_ptrs, c, mask=mask, cache_modifier=".wt")
@triton.jit
def _splitk_reduce_kernel(
src_ptr,
dst_ptr,
m,
n,
actual_ksplit,
stride_sk,
stride_sm,
stride_sn,
stride_dm,
stride_dn,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
MAX_KSPLIT: tl.constexpr,
):
pid_m = tl.program_id(axis=0)
pid_n = tl.program_id(axis=1)
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, MAX_KSPLIT)
src_ptrs = (
src_ptr
+ offs_k[:, None, None] * stride_sk
+ offs_m[None, :, None] * stride_sm
+ offs_n[None, None, :] * stride_sn
)
src_mask = (
(offs_k[:, None, None] < actual_ksplit)
& (offs_m[None, :, None] < m)
& (offs_n[None, None, :] < n)
)
vals = tl.load(src_ptrs, mask=src_mask, other=0.0)
out = tl.sum(vals, axis=0).to(dst_ptr.type.element_ty)
dst_ptrs = dst_ptr + offs_m[:, None] * stride_dm + offs_n[None, :] * stride_dn
tl.store(dst_ptrs, out, mask=(offs_m[:, None] < m) & (offs_n[None, :] < n))
def _as_u8(x: torch.Tensor) -> torch.Tensor:
y = x.contiguous()
if y.dtype != torch.uint8:
y = y.view(torch.uint8)
return y
def _pad_scale_cols(x: torch.Tensor) -> torch.Tensor:
y = _as_u8(x)
rows, cols = y.shape
cols_pad = ((cols + 7) // 8) * 8
if cols_pad == cols:
return y
out = torch.zeros((rows, cols_pad), dtype=torch.uint8, device=y.device)
out[:, :cols] = y
return out
def _shuffle_e8m0(x: torch.Tensor) -> torch.Tensor:
y = _pad_scale_cols(x)
rows, cols = y.shape
rows_pad = ((rows + 255) // 256) * 256
cols_pad = ((cols + 7) // 8) * 8
out = torch.zeros((rows_pad, cols_pad), dtype=torch.uint8, device=y.device)
out[:rows, :cols] = y
return (
out.view(rows_pad // 32, 2, 16, cols_pad // 8, 2, 4)
.permute(0, 3, 5, 2, 4, 1)
.contiguous()
.view(rows_pad, cols_pad)
)
def _view_preshuffle_weight(x: torch.Tensor) -> torch.Tensor:
y = _as_u8(x)
rows, cols = y.shape
return y.view(rows // 16, cols * 16)
def _view_preshuffle_scale(x: torch.Tensor) -> torch.Tensor:
y = _as_u8(x)
rows, cols = y.shape
return y.view(rows // 32, cols * 32)
def _prepare_a_scale(x: torch.Tensor, m: int) -> tuple[torch.Tensor, bool]:
if m < 32:
return _pad_scale_cols(x), True
return _view_preshuffle_scale(_shuffle_e8m0(x)), False
_RUNTIME = {
"torch": torch,
"triton": triton,
"dynamic_mxfp4_quant": dynamic_mxfp4_quant,
"kernel": _mxfp4_preshuffle_kernel,
"reduce_kernel": _splitk_reduce_kernel,
"as_u8": _as_u8,
"view_preshuffle_weight": _view_preshuffle_weight,
"view_preshuffle_scale": _view_preshuffle_scale,
"prepare_a_scale": _prepare_a_scale,
}
return _RUNTIME
def custom_kernel(data: tuple[Any, ...]) -> Any:
rt = _load_runtime()
torch = rt["torch"]
triton = rt["triton"]
a, _b, _b_q, b_shuffle, b_scale_sh = data
if a.ndim != 2:
raise RuntimeError(f"A must be 2D, got {tuple(a.shape)}")
m, k = a.shape
n = b_shuffle.shape[0]
a_q, a_scale = rt["dynamic_mxfp4_quant"](a.contiguous())
a_q_u8 = rt["as_u8"](a_q)
a_scale_view, small_m = rt["prepare_a_scale"](a_scale, m)
b_view = rt["view_preshuffle_weight"](b_shuffle)
b_scale_view = rt["view_preshuffle_scale"](b_scale_sh)
(
block_m,
block_n,
block_k,
group_m,
num_warps,
num_stages,
waves_per_eu,
num_ksplit,
even_k,
cache_hint,
) = _pick_config(m, n, k)
grid = (triton.cdiv(m, block_m) * triton.cdiv(n, block_n) * num_ksplit,)
if num_ksplit == 1:
out = torch.empty((m, n), dtype=torch.bfloat16, device=a.device)
partial = out
stride_ck = 0
else:
partial = torch.empty((num_ksplit, m, n), dtype=torch.float32, device=a.device)
out = torch.empty((m, n), dtype=torch.bfloat16, device=a.device)
stride_ck = partial.stride(0)
rt["kernel"][grid](
a_q_u8,
b_view,
partial,
a_scale_view,
b_scale_view,
m,
n,
a_q_u8.shape[1],
b_view.shape[0],
a_scale_view.shape[0],
b_scale_view.shape[0],
a_q_u8.stride(0),
a_q_u8.stride(1),
b_view.stride(0),
b_view.stride(1),
stride_ck,
partial.stride(-2),
partial.stride(-1),
a_scale_view.stride(0),
a_scale_view.stride(1),
b_scale_view.stride(0),
b_scale_view.stride(1),
BLOCK_M=block_m,
BLOCK_N=block_n,
BLOCK_K=block_k,
GROUP_SIZE_M=group_m,
SMALL_M=small_m,
NUM_KSPLIT=num_ksplit,
SPLITK_BLOCK_SIZE=(k + num_ksplit - 1) // num_ksplit,
EVEN_K=even_k,
CACHE_HINT=cache_hint,
num_warps=num_warps,
num_stages=num_stages,
waves_per_eu=waves_per_eu,
matrix_instr_nonkdim=16,
)
if num_ksplit == 1:
return out
reduce_block_m = 16
reduce_block_n = 64
reduce_grid = (triton.cdiv(m, reduce_block_m), triton.cdiv(n, reduce_block_n))
max_ksplit = 1 << (num_ksplit - 1).bit_length()
rt["reduce_kernel"][reduce_grid](
partial,
out,
m,
n,
num_ksplit,
partial.stride(0),
partial.stride(1),
partial.stride(2),
out.stride(0),
out.stride(1),
BLOCK_M=reduce_block_m,
BLOCK_N=reduce_block_n,
MAX_KSPLIT=max_ksplit,
num_warps=4,
num_stages=1,
)
return out
scrolls · 512 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