submission 670896
zaiji100 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 228 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-670896?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:9d3d966a43d739097367380fd9287832ebe9e32fb787d9f6d3821bfc5fd10935
license declaredunknown
license concludedunknown
authorszaiji100
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
stages = 1
num_warps=meta['nw'], waves_per_eu=0, num_stages=1)Kernel source
submission.py228 lines
#!POPCORN leaderboard amd-mxfp4-mm
#!POPCORN gpu MI355X
from task import input_t, output_t
import os
from typing import Dict, List, Optional, Set, Tuple
os.environ.setdefault("PYTORCH_ROCM_ARCH", "gfx950")
os.environ.setdefault("CXX", "clang++")
os.environ.setdefault("AITER_KSPLIT", "1")
import torch
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
_fused_mod = None
# ---- Triton fused quant+shuffle (proven fallback) ----
_triton_ok = False
try:
import triton
import triton.language as tl
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
_triton_ok = True
except Exception:
pass
_gemm_asm = getattr(aiter, "gemm_a4w4_asm", None)
_gemm_a4w4 = aiter.gemm_a4w4
_get_padded_m = getattr(aiter, "get_padded_m", None)
if _triton_ok:
@triton.heuristics({
"EVEN_M_N": lambda args: args["M"] % args["BLOCK_SIZE_M"] == 0
and args["N"] % (args["BLOCK_SIZE_N"] * args["NUM_ITER"]) == 0,
})
@triton.jit
def _fused_quant_shuffle_kernel(
x_ptr, x_fp4_ptr, bs_ptr,
stride_x_m_in, stride_x_n_in, stride_x_fp4_m_in, stride_x_fp4_n_in,
M, N,
SCALE_N: tl.constexpr, SCALE_M_PAD: tl.constexpr, SCALE_N_PAD: tl.constexpr,
BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr,
NUM_ITER: tl.constexpr, NUM_STAGES: tl.constexpr,
MXFP4_QUANT_BLOCK_SIZE: tl.constexpr, EVEN_M_N: tl.constexpr,
):
pid_m = tl.program_id(0)
start_n = tl.program_id(1) * NUM_ITER
stride_x_m = tl.cast(stride_x_m_in, tl.int64)
stride_x_n = tl.cast(stride_x_n_in, tl.int64)
stride_x_fp4_m = tl.cast(stride_x_fp4_m_in, tl.int64)
stride_x_fp4_n = tl.cast(stride_x_fp4_n_in, tl.int64)
NUM_QUANT_BLOCKS: tl.constexpr = BLOCK_SIZE_N // MXFP4_QUANT_BLOCK_SIZE
for pid_n in tl.range(start_n, min(start_n + NUM_ITER, N), num_stages=NUM_STAGES):
x_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
x_offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N)
x_offs = x_offs_m[:, None] * stride_x_m + x_offs_n[None, :] * stride_x_n
if EVEN_M_N:
x = tl.load(x_ptr + x_offs, cache_modifier=".cg").to(tl.float32)
else:
x_mask = (x_offs_m < M)[:, None] & (x_offs_n < N)[None, :]
x = tl.load(x_ptr + x_offs, mask=x_mask, cache_modifier=".cg").to(tl.float32)
out_tensor, bs_e8m0 = _mxfp4_quant_op(x, BLOCK_SIZE_N, BLOCK_SIZE_M, MXFP4_QUANT_BLOCK_SIZE)
out_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
out_offs_n = pid_n * BLOCK_SIZE_N // 2 + tl.arange(0, BLOCK_SIZE_N // 2)
out_offs = out_offs_m[:, None] * stride_x_fp4_m + out_offs_n[None, :] * stride_x_fp4_n
if EVEN_M_N:
tl.store(x_fp4_ptr + out_offs, out_tensor)
else:
tl.store(x_fp4_ptr + out_offs, out_tensor, mask=(out_offs_m < M)[:, None] & (out_offs_n < (N // 2))[None, :])
bs_offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)
bs_offs_n = pid_n * NUM_QUANT_BLOCKS + tl.arange(0, NUM_QUANT_BLOCKS)
num_bs_cols = (N + MXFP4_QUANT_BLOCK_SIZE - 1) // MXFP4_QUANT_BLOCK_SIZE
d0 = bs_offs_m[:, None] // 32
d1_full = bs_offs_m[:, None] % 32
d2 = d1_full % 16
d1 = d1_full // 16
d3 = bs_offs_n[None, :] // 8
d4_full = bs_offs_n[None, :] % 8
d5 = d4_full % 4
d4 = d4_full // 4
bs_shuffled_offs = d1 + d4 * 2 + d2 * 4 + d5 * 64 + d3 * 256 + d0 * 32 * SCALE_N_PAD
bs_valid = (bs_offs_m < M)[:, None] & (bs_offs_n < num_bs_cols)[None, :]
bs_e8m0 = tl.where(bs_valid, bs_e8m0, 127)
bs_pad = (bs_offs_m < SCALE_M_PAD)[:, None] & (bs_offs_n < SCALE_N_PAD)[None, :]
tl.store(bs_ptr + bs_shuffled_offs, bs_e8m0, mask=bs_pad)
def _mk(tile_m, tile_n=128):
s = f"f4gemm_bf16_per1x32Fp4_BpreShuffle_{tile_m}x{tile_n}"
return f"_ZN5aiter{len(s)}{s}E"
_K32 = _mk(32)
_ASM_CANDIDATES: Dict[Tuple[int, int, int], List[Tuple[str, int]]] = {
(4, 2880, 512): [(_K32, 0)],
(8, 2112, 7168): [(_K32, 0), (_K32, 2)],
(16, 2112, 7168): [(_K32, 0), (_K32, 2)],
(16, 3072, 1536): [(_K32, 0)],
(32, 2880, 512): [(_K32, 0)],
(32, 4096, 512): [(_K32, 0)],
}
_quant_cache: Dict[Tuple[int, int, int], Tuple[torch.Tensor, torch.Tensor, dict]] = {}
_good_cfg: Dict[Tuple[int, int, int], Tuple[str, int]] = {}
_bad_cfg: Set[Tuple[int, int, int, str, int]] = set()
_padded_rows: Dict[Tuple[int, int, int], int] = {}
_out_bufs: Dict[Tuple[int, int, int], torch.Tensor] = {}
_fused_ok: Optional[bool] = None
_b_raw_cache: Dict[int, torch.Tensor] = {}
_bsc_raw_cache: Dict[int, torch.Tensor] = {}
def _get_raw_bsc(b_sc, N, K):
ptr = b_sc.data_ptr()
c = _bsc_raw_cache.get(ptr)
if c is not None: return c
s = b_sc.view(torch.uint8)
sm, sn = s.shape
s = s.view(sm//32, sn//8, 4, 16, 2, 2).permute(0,5,3,1,4,2).contiguous().view(sm, sn)
raw = s[:N, :(K+31)//32].contiguous()
_bsc_raw_cache[ptr] = raw
return raw
def _ensure_quant(dev_idx, M, K, device):
key = (dev_idx, M, K)
cached = _quant_cache.get(key)
if cached is not None: return cached
sn_valid = (K + 31) // 32
sn_pad = ((sn_valid + 7) // 8) * 8
sm_pad = ((M + 255) // 256) * 256
if M <= 32:
bm, bn, ni, nw, ns = triton.next_power_of_2(M), 32, 1, 1, 1
else:
ni, bm, bn, nw, ns = 4, 64, 64, 4, 2
if K <= 16384: bm, bn = 32, 128
if K <= 1024:
ni, ns, nw = 1, 1, 4
bn = max(32, min(256, triton.next_power_of_2(K)))
bm = min(8, triton.next_power_of_2(M))
if M <= 32 and K > 1024:
bm, bn, ni, nw, ns = triton.next_power_of_2(M), 128, 1, 4, 1
grid = (triton.cdiv(M, bm), triton.cdiv(K, bn * ni))
x_fp4 = torch.empty((M, K // 2), dtype=torch.uint8, device=device)
bs = torch.empty((sm_pad, sn_pad), dtype=torch.uint8, device=device)
meta = {'sn_valid': sn_valid, 'sn_pad': sn_pad, 'sm_pad': sm_pad,
'bm': bm, 'bn': bn, 'ni': ni, 'nw': nw, 'ns': ns, 'grid': grid}
_quant_cache[key] = (x_fp4, bs, meta)
return (x_fp4, bs, meta)
def _get_padded(m, n, k):
key = (m, n, k)
r = _padded_rows.get(key)
if r is not None: return r
r = m
if _get_padded_m:
try: r = int(_get_padded_m(m, n, k, 32))
except: pass
_padded_rows[key] = r
return r
@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
a = data[0]
b_sh = data[3]
b_sc = data[4]
if not a.is_contiguous():
a = a.contiguous()
M, K = a.shape
N = b_sh.shape[0]
key = (M, N, K)
# ---- Triton fused quant+shuffle + ASM GEMM ----
global _fused_ok
if _triton_ok and _fused_ok is not False:
try:
x_fp4, bs, meta = _ensure_quant(a.device.index or 0, M, K, a.device)
_fused_quant_shuffle_kernel[meta['grid']](
a, x_fp4, bs, K, 1, K >> 1, 1,
M=M, N=K,
SCALE_N=meta['sn_valid'], SCALE_M_PAD=meta['sm_pad'], SCALE_N_PAD=meta['sn_pad'],
BLOCK_SIZE_M=meta['bm'], BLOCK_SIZE_N=meta['bn'],
NUM_ITER=meta['ni'], NUM_STAGES=meta['ns'],
MXFP4_QUANT_BLOCK_SIZE=32,
num_warps=meta['nw'], waves_per_eu=0, num_stages=1)
aq = x_fp4.view(dtypes.fp4x2)
a_sc = bs.view(dtypes.fp8_e8m0)
_fused_ok = True
except Exception:
_fused_ok = False
xq, bse = dynamic_mxfp4_quant(a)
aq = xq.view(dtypes.fp4x2)
a_sc = e8m0_shuffle(bse).view(dtypes.fp8_e8m0)
else:
xq, bse = dynamic_mxfp4_quant(a)
aq = xq.view(dtypes.fp4x2)
a_sc = e8m0_shuffle(bse).view(dtypes.fp8_e8m0)
m, n, k = M, N, K
if _gemm_asm is not None:
cfg = _good_cfg.get(key)
if cfg is None:
for kn, sk in _ASM_CANDIDATES.get(key, []):
if (m, n, k, kn, sk) not in _bad_cfg:
cfg = (kn, sk)
break
if cfg is not None:
kn, sk = cfg
pr = _get_padded(m, n, k)
buf_key = (aq.device.index or 0, pr, n)
out = _out_bufs.get(buf_key)
if out is None:
out = torch.empty((pr, n), dtype=torch.bfloat16, device=aq.device)
_out_bufs[buf_key] = out
try:
if sk > 0: out.zero_()
_gemm_asm(aq.view(m, -1), b_sh, a_sc.view(m, -1), b_sc,
out, kn, None, 1.0, 0.0, True, sk)
_good_cfg[key] = cfg
return out[:m]
except Exception:
_bad_cfg.add((m, n, k, kn, sk))
return _gemm_a4w4(aq, b_sh, a_sc, b_sc, dtype=dtypes.bf16, bpreshuffle=True)
scrolls · 228 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