submission 627274
Bortlesboat · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 342 lines, June 9 Researcher Reciprocity License v1.0.
v78_splitk0.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-627274?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:f50d65705980d918670afdd84686a443d0865fca357cd03cebfc093a68b30301
license declaredunknown
license concludedunknown
authorsBortlesboat
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
autotune
@triton.autotune(num-warps = 4
triton.Config({'BSM': 4, 'BSN': 512, 'NI': 1, 'NS': 1}, num_warps=4),split-k
0, out.stride(0), out.stride(1), # stride_ck=0 for no split-KKernel source
v78_splitk0.py342 lines
"""v53: Use aiter's _gemm_a16wfp4_preshuffle_kernel directly with PREQUANT=True.
Based on dgavriloff/amd-structkernel approach (8.667μs proven).
Key insight: the preshuffle kernel handles B_shuffle layout natively with
coalesced loads + compile-time reshape/permute/trans. No manual B transposition needed.
Shapes 1-5 (M<=64): _gemm_a16wfp4_preshuffle_kernel with PREQUANT=True
Shape 6 (M=256): quant A + ASM gemm_a4w4_asm (32x128 tile)
"""
from task import input_t, output_t
import os, sys
os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
os.environ["HSA_TOOLS_LIB"] = ""
import torch
import triton
import triton.language as tl
from aiter import dtypes
import aiter
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
_fp4x2 = dtypes.fp4x2; _fp8_e8m0 = dtypes.fp8_e8m0; _bf16 = dtypes.bf16
# Try importing aiter's preshuffle kernel
_HAS_PRESHUFFLE = False
try:
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import (
_gemm_a16wfp4_preshuffle_kernel,
)
_HAS_PRESHUFFLE = True
except ImportError:
print("[v53] WARN: _gemm_a16wfp4_preshuffle_kernel not available", file=sys.stderr)
# Try importing gluon reduce kernel for split-K
_HAS_GLUON_REDUCE = False
try:
from aiter.ops.triton.gluon.gemm_afp4wfp4 import (
_gemm_afp4wfp4_reduce_kernel as _gluon_reduce_kernel,
)
_HAS_GLUON_REDUCE = True
except ImportError:
pass
if not _HAS_GLUON_REDUCE:
try:
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
_gemm_afp4wfp4_reduce_kernel as _gluon_reduce_kernel,
)
_HAS_GLUON_REDUCE = True
except ImportError:
print("[v53] WARN: reduce kernel not available", file=sys.stderr)
# Try importing _mxfp4_quant_op for fallback fused quant
try:
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
_HAS_IQ = True
except ImportError:
_HAS_IQ = False
# ---- Shape-specific configs (from dgavriloff's proven v224) ----
# (M, N, K) -> (BSM, BSN, BSK, warps, stages, waves_per_eu, cache, split_k)
SHAPE_CONFIGS = {
(4, 2880, 512): (4, 128, 256, 4, 2, 0, ".cg", 1),
(16, 2112, 7168): (8, 128, 256, 4, 2, 2, ".cg", 7),
(32, 4096, 512): (8, 128, 256, 4, 2, 2, None, 1),
(32, 2880, 512): (8, 128, 256, 4, 2, 2, None, 1),
(64, 7168, 2048): (16, 128, 256, 4, 2, 2, ".cg", 1),
}
# ---- Buffer caches ----
_ob = {}; _ws = {}; _bw = {}; _bsc = {}; _fo = {}
def _get_obuf(M, N, d):
k = (M, N)
if k not in _ob:
_ob[k] = torch.empty((M, N), dtype=torch.bfloat16, device=d)
return _ob[k]
def _reshape_b(B_shuffle, N, K):
"""Reshape B_shuffle for preshuffle kernel: (N//16, (K//2)*16)"""
k = (N, K)
if k in _bw:
bw, ref = _bw[k]
if ref.data_ptr() == B_shuffle.data_ptr():
return bw
bw = B_shuffle.view(torch.uint8).reshape(N // 16, (K // 2) * 16)
_bw[k] = (bw, B_shuffle)
return bw
def _reshape_bsc(B_scale_sh, N, K):
"""Reshape B_scale_sh for preshuffle kernel: (bs0//32, bs1*32)"""
k = (N, K)
if k in _bsc:
bsc, ref = _bsc[k]
if ref.data_ptr() == B_scale_sh.data_ptr():
return bsc
bs = B_scale_sh.view(torch.uint8)
bs0, bs1 = bs.shape
bsc = bs.reshape(bs0 // 32, bs1 * 32)
_bsc[k] = (bsc, B_scale_sh)
return bsc
def _run_preshuffle(A, B_shuffle, B_scale_sh, M, N, K, cfg):
BSM, BSN, BSK, warps, stages, wpe, cache, num_ksplit = cfg
B_w = _reshape_b(B_shuffle, N, K)
B_sc = _reshape_bsc(B_scale_sh, N, K)
K_kernel = K // 2
if num_ksplit == 1:
out = _get_obuf(M, N, A.device)
grid_size = triton.cdiv(M, BSM) * triton.cdiv(N, BSN)
_gemm_a16wfp4_preshuffle_kernel[(grid_size,)](
A, B_w, out, B_sc,
M, N, K_kernel,
A.stride(0), A.stride(1),
B_w.stride(0), B_w.stride(1),
0, out.stride(0), out.stride(1), # stride_ck=0 for no split-K
B_sc.stride(0), B_sc.stride(1),
BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=BSN, BLOCK_SIZE_K=BSK,
GROUP_SIZE_M=1,
NUM_KSPLIT=1,
SPLITK_BLOCK_SIZE=2 * K_kernel,
num_warps=warps, num_stages=stages, waves_per_eu=wpe,
matrix_instr_nonkdim=16,
PREQUANT=True,
cache_modifier=cache,
)
return out
else:
# Split-K path
wk = (num_ksplit, M, N)
if wk not in _ws:
_ws[wk] = torch.empty(wk, dtype=torch.float32, device=A.device)
y_pp = _ws[wk]
# Compute SPLITK_BLOCK_SIZE: each split handles K_kernel/num_ksplit elements
# SPLITK_BLOCK_SIZE = ceil(K_kernel / num_ksplit) rounded up to BSK
k_per_split = (K_kernel + num_ksplit - 1) // num_ksplit
SPLITK_BLOCK_SIZE = ((k_per_split + BSK - 1) // BSK) * BSK * 2
grid_size = num_ksplit * triton.cdiv(M, BSM) * triton.cdiv(N, BSN)
_gemm_a16wfp4_preshuffle_kernel[(grid_size,)](
A, B_w, y_pp, B_sc,
M, N, K_kernel,
A.stride(0), A.stride(1),
B_w.stride(0), B_w.stride(1),
y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
B_sc.stride(0), B_sc.stride(1),
BLOCK_SIZE_M=BSM, BLOCK_SIZE_N=BSN, BLOCK_SIZE_K=BSK,
GROUP_SIZE_M=1,
NUM_KSPLIT=num_ksplit,
SPLITK_BLOCK_SIZE=SPLITK_BLOCK_SIZE,
num_warps=warps, num_stages=stages, waves_per_eu=wpe,
matrix_instr_nonkdim=16,
PREQUANT=True,
cache_modifier=cache,
)
# Reduce split-K partials
out = _get_obuf(M, N, A.device)
ACTUAL_KSPLIT = triton.cdiv(K_kernel, SPLITK_BLOCK_SIZE // 2)
if _HAS_GLUON_REDUCE:
reduce_grid = (triton.cdiv(M, 16), triton.cdiv(N, 64))
_gluon_reduce_kernel[reduce_grid](
y_pp, out,
M, N,
y_pp.stride(0), y_pp.stride(1), y_pp.stride(2),
out.stride(0), out.stride(1),
16, 64,
ACTUAL_KSPLIT,
triton.next_power_of_2(num_ksplit),
)
else:
# Fallback: simple sum in Python
out.copy_(y_pp.sum(dim=0).to(torch.bfloat16))
return out
# ---- A quant for ASM path (M=256) ----
@triton.autotune(
configs=[
triton.Config({'BSM': 4, 'BSN': 512, 'NI': 1, 'NS': 1}, num_warps=4),
triton.Config({'BSM': 8, 'BSN': 256, 'NI': 1, 'NS': 1}, num_warps=4),
triton.Config({'BSM': 8, 'BSN': 512, 'NI': 1, 'NS': 1}, num_warps=4),
triton.Config({'BSM': 16, 'BSN': 128, 'NI': 4, 'NS': 2}, num_warps=4),
triton.Config({'BSM': 32, 'BSN': 512, 'NI': 1, 'NS': 1}, num_warps=4),
triton.Config({'BSM': 64, 'BSN': 128, 'NI': 4, 'NS': 2}, num_warps=4),
],
key=['M', 'N'],
)
@triton.jit
def _quant_kernel(x, fp, bs, sxm, sxn, sfm, sfn, M, N, snp,
BSM: tl.constexpr, BSN: tl.constexpr,
NI: tl.constexpr, NS: tl.constexpr, QBS: tl.constexpr):
pm = tl.program_id(0); sn = tl.program_id(1) * NI
xm = tl.cast(sxm, tl.int64); xn = tl.cast(sxn, tl.int64)
fm = tl.cast(sfm, tl.int64); fn = tl.cast(sfn, tl.int64)
NQB: tl.constexpr = BSN // QBS
for pn in tl.range(sn, min(sn + NI, N), num_stages=NS):
om = pm * BSM + tl.arange(0, BSM)
on = pn * BSN + tl.arange(0, BSN)
v = tl.load(x + om[:, None] * xm + on[None, :] * xn,
mask=(om < M)[:, None] & (on < N)[None, :], other=0.0,
cache_modifier=".cg").to(tl.float32)
f4, sc = _mxfp4_quant_op(v, BSN, BSM, QBS)
fo = pm * BSM + tl.arange(0, BSM)
fno = pn * BSN // 2 + tl.arange(0, BSN // 2)
tl.store(fp + fo[:, None] * fm + fno[None, :] * fn, f4,
mask=(fo < M)[:, None] & (fno < N // 2)[None, :])
sm = pm * BSM + tl.arange(0, BSM)
sk = pn * NQB + tl.arange(0, NQB)
d0 = sm // 32; d1 = (sm % 32) // 16; d2 = sm % 16
d3 = sk // 8; d4 = (sk % 8) // 4; d5 = sk % 4
fl = ((d0 * (snp * 32))[:, None] + (d3 * 256)[None, :]
+ (d5 * 64)[None, :] + (d2 * 4)[:, None]
+ (d4 * 2)[None, :] + d1[:, None])
tl.store(bs + fl, sc,
mask=(sm < M)[:, None] & (sk < N // QBS)[None, :])
_qb = {}
def _get_qbuf(M, K, d):
k = (M, K)
if k not in _qb:
nq = K // 32; sp = ((M + 255) // 256) * 256; snp = ((nq + 7) // 8) * 8
_qb[k] = (torch.empty((M, K // 2), dtype=torch.uint8, device=d),
torch.zeros(sp * snp, dtype=torch.uint8, device=d), sp, snp)
return _qb[k]
def _get_snp(K):
return ((K // 32 + 7) // 8) * 8
def _quant_a(A):
M, K = A.shape
f, b, sp, snp = _get_qbuf(M, K, A.device)
g = lambda m: (triton.cdiv(M, m['BSM']), triton.cdiv(K, m['BSN'] * m['NI']))
_quant_kernel[g](A, f, b, A.stride(0), A.stride(1),
f.stride(0), f.stride(1), M, K, snp, QBS=32)
return f.view(_fp4x2), b.view(sp, snp).view(_fp8_e8m0)
def _kn(t):
i = f"f4gemm_bf16_per1x32Fp4_BpreShuffle_{t}"
return f"_ZN5aiter{len(i)}{i}E"
_ga = None; _gg = None; _ao = {}
def custom_kernel(data: input_t) -> output_t:
global _ga, _gg
A, B, Bq, Bs, Bss = data # Bs=B_shuffle, Bss=B_scale_sh
A = A.contiguous()
M, K = A.shape; N = B.shape[0]
key = (M, N, K)
# ---- PRESHUFFLE KERNEL: M<=64 (shapes 1-5) ----
if _HAS_PRESHUFFLE and M <= 64:
cfg = SHAPE_CONFIGS.get(key)
if cfg is None:
# Default config for unknown shapes
if K > 2048:
cfg = (8, 128, 256, 4, 2, 2, ".cg", max(1, K // 1024))
else:
cfg = (min(M, 16), 128, 256, 4, 2, 2, ".cg", 1)
pk = ("ps", key)
if pk not in _fo:
try:
C = _run_preshuffle(A, Bs, Bss, M, N, K, cfg)
# Verify correctness on first call
if _gg is None: _gg = aiter.gemm_a4w4
Aq, Ac = _quant_a(A) if _HAS_IQ else _sq(A)
ref = _gg(Aq, Bs, Ac, Bss, dtype=_bf16, bpreshuffle=True)
me = (C.float() - ref.float()).abs().max().item()
_fo[pk] = me < 2.0
if _fo[pk]:
print(f"[v53] PRESHUFFLE OK {key} err={me:.1f}", file=sys.stderr)
return C
else:
print(f"[v53] PRESHUFFLE FAIL {key} err={me:.1f}", file=sys.stderr)
except Exception as e:
_fo[pk] = False
print(f"[v53] PRESHUFFLE ERR {key}: {e}", file=sys.stderr)
elif _fo[pk]:
return _run_preshuffle(A, Bs, Bss, M, N, K, cfg)
# ---- M=256: quant A + ASM GEMM (two-phase) ----
Aq, Ac = _quant_a(A) if _HAS_IQ else _sq(A)
if _gg is None: _gg = aiter.gemm_a4w4
if _ga is None: _ga = aiter.gemm_a4w4_asm
if M == 256 and K >= 1536:
tile = "32x128"
sp = None # no split-K (aiter get_GEMM_config recommends splitK=0)
tk = ("asm256", tile)
if tk not in _ao:
try:
out = _get_obuf(M, N, A.device)
_ga(Aq, Bs, Ac, Bss, out, _kn(tile), bpreshuffle=True, log2_k_split=sp)
_ao[tk] = True
except Exception:
_ao[tk] = False
if _ao.get(tk, False):
out = _get_obuf(M, N, A.device)
return _ga(Aq, Bs, Ac, Bss, out, _kn(tile), bpreshuffle=True, log2_k_split=sp)
if M == 64:
tile = "32x128"
sp = None # try no split-K
tk = ("asm64", tile)
if tk not in _ao:
try:
out = _get_obuf(M, N, A.device)
_ga(Aq, Bs, Ac, Bss, out, _kn(tile), bpreshuffle=True, log2_k_split=sp)
_ao[tk] = True
except Exception:
_ao[tk] = False
if _ao.get(tk, False):
out = _get_obuf(M, N, A.device)
return _ga(Aq, Bs, Ac, Bss, out, _kn(tile), bpreshuffle=True, log2_k_split=sp)
if M <= 32:
tile = "32x128"
sp = 1 if K >= 2048 else None
if tile not in _ao:
try:
out = _get_obuf(M, N, A.device)
_ga(Aq, Bs, Ac, Bss, out, _kn(tile), bpreshuffle=True, log2_k_split=sp)
_ao[tile] = True
except Exception:
_ao[tile] = False
if _ao.get(tile, False):
out = _get_obuf(M, N, A.device)
return _ga(Aq, Bs, Ac, Bss, out, _kn(tile), bpreshuffle=True, log2_k_split=sp)
# ---- FALLBACK: gemm_a4w4 ----
return _gg(Aq, Bs, Ac, Bss, dtype=_bf16, bpreshuffle=True)
def _sq(A):
Aq, Ac = dynamic_mxfp4_quant(A)
Ac = e8m0_shuffle(Ac)
return Aq.view(_fp4x2), Ac.view(_fp8_e8m0)
scrolls · 342 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