submission 748790
gogogo_666 · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 351 lines, June 9 Researcher Reciprocity License v1.0.
submission_attempt38_shape2_splitk7.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-748790?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:a005e4975a050ceea4e7f25ebffb621954169e94d0cf61ba31e4b04b916e75fc
license declaredunknown
license concludedunknown
authorsgogogo_666
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
split-k
Attempt 38: shape-2 no-graph Triton with splitK=7.Kernel source
submission_attempt38_shape2_splitk7.py351 lines
"""
Attempt 38: shape-2 no-graph Triton with splitK=7.
Rationale:
- Attempt 32 remains the best shape-2-safe point, but it still uses `NUM_KSPLIT=14`
and pays a non-trivial reduce over 14 partial tiles.
- Recent metadata-only probes have all regressed, so the next credible lever is a
structural splitK change on the same kernel family.
- Setting shape-2 `NUM_KSPLIT=7` halves the number of partials and reduce work while
preserving the same no-graph preshuffle kernel path.
- Other benchmark shapes stay on the locked best wrapper path.
"""
from task import input_t, output_t
import os
import torch
import triton
import triton.language as tl
_CUSTOM_CFG_TEXT = """cu_num,M,N,K,kernelId,splitK,us,kernelName,tflops,bw,errRatio
256,4,2880,512,21,0,4.3600,_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E,0,0,0
256,16,2112,7168,21,0,12.8400,_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128E,0,0,0
256,32,4096,512,29,0,4.5800,_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E,0,0,0
256,32,2880,512,29,0,4.5800,_ZN5aiter41f4gemm_bf16_per1x32Fp4_BpreShuffle_64x128E,0,0,0
"""
_CUSTOM_CFG_PATH = "/tmp/aiter_config_a4w4_coco_locked_v3.csv"
with open(_CUSTOM_CFG_PATH, "w", encoding="ascii") as _f:
_f.write(_CUSTOM_CFG_TEXT)
os.environ["AITER_CONFIG_GEMM_A4W4"] = (
f"{_CUSTOM_CFG_PATH}{os.pathsep}/home/runner/aiter/aiter/configs/a4w4_blockscale_tuned_gemm.csv"
)
from aiter.ops.triton._triton_kernels.quant.quant import _mxfp4_quant_op
from aiter.ops.triton.utils._triton.pid_preprocessing import pid_grid
@triton.jit
def _fixed_preshuffle_kernel(
a_ptr,
b_ptr,
c_ptr,
b_scales_ptr,
M,
N,
K,
stride_am,
stride_ak,
stride_bn,
stride_bk,
stride_ck,
stride_cm,
stride_cn,
stride_bsn,
stride_bsk,
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr,
GROUP_SIZE_M: tl.constexpr,
NUM_KSPLIT: tl.constexpr,
SPLITK_BLOCK_SIZE: tl.constexpr,
num_warps: tl.constexpr,
num_stages: tl.constexpr,
waves_per_eu: tl.constexpr,
matrix_instr_nonkdim: tl.constexpr,
cache_modifier: tl.constexpr,
):
pid_unified = tl.program_id(axis=0)
pid_k = pid_unified % NUM_KSPLIT
pid = pid_unified // NUM_KSPLIT
num_pid_m = tl.cdiv(M, BLOCK_SIZE_M)
num_pid_n = tl.cdiv(N, BLOCK_SIZE_N)
if NUM_KSPLIT == 1:
pid_m, pid_n = pid_grid(pid, num_pid_m, num_pid_n, GROUP_SIZE_M=GROUP_SIZE_M)
else:
pid_m = pid // num_pid_n
pid_n = pid % num_pid_n
SCALE_GROUP_SIZE: tl.constexpr = 32
if (pid_k * SPLITK_BLOCK_SIZE // 2) < K:
num_k_iter = tl.cdiv(SPLITK_BLOCK_SIZE // 2, BLOCK_SIZE_K // 2)
offs_k_bf16 = tl.arange(0, BLOCK_SIZE_K)
offs_k_split_bf16 = pid_k * SPLITK_BLOCK_SIZE + offs_k_bf16
offs_am = (pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M)) % M
a_ptrs = a_ptr + (
offs_am[:, None] * stride_am + offs_k_split_bf16[None, :] * stride_ak
)
offs_k_shuffle_arr = tl.arange(0, (BLOCK_SIZE_K // 2) * 16)
offs_k_shuffle = pid_k * (SPLITK_BLOCK_SIZE // 2) * 16 + offs_k_shuffle_arr
offs_bn = (pid_n * (BLOCK_SIZE_N // 16) + tl.arange(0, BLOCK_SIZE_N // 16)) % N
b_ptrs = b_ptr + (
offs_bn[:, None] * stride_bn + offs_k_shuffle[None, :] * stride_bk
)
offs_bsn = (
pid_n * (BLOCK_SIZE_N // 32) + tl.arange(0, (BLOCK_SIZE_N // 32))
) % N
offs_ks = (pid_k * (SPLITK_BLOCK_SIZE // SCALE_GROUP_SIZE) * 32) + tl.arange(
0, BLOCK_SIZE_K // SCALE_GROUP_SIZE * 32
)
b_scale_ptrs = (
b_scales_ptr
+ offs_bsn[:, None] * stride_bsn
+ offs_ks[None, :] * stride_bsk
)
accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32)
for _ in range(pid_k * num_k_iter, (pid_k + 1) * num_k_iter):
b_scales = (
tl.load(b_scale_ptrs, cache_modifier=cache_modifier)
.reshape(BLOCK_SIZE_N // 32, BLOCK_SIZE_K // SCALE_GROUP_SIZE // 8, 4, 16, 2, 2, 1)
.permute(0, 5, 3, 1, 4, 2, 6)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // SCALE_GROUP_SIZE)
)
a_bf16 = tl.load(a_ptrs)
b = tl.load(b_ptrs, cache_modifier=cache_modifier)
b = (
b.reshape(1, BLOCK_SIZE_N // 16, BLOCK_SIZE_K // 64, 2, 16, 16)
.permute(0, 1, 4, 2, 3, 5)
.reshape(BLOCK_SIZE_N, BLOCK_SIZE_K // 2)
.trans(1, 0)
)
a, a_scales = _mxfp4_quant_op(a_bf16, BLOCK_SIZE_K, BLOCK_SIZE_M, 32)
accumulator += tl.dot_scaled(a, a_scales, "e2m1", b, b_scales, "e2m1")
a_ptrs += BLOCK_SIZE_K * stride_ak
b_ptrs += (BLOCK_SIZE_K // 2) * 16 * stride_bk
b_scale_ptrs += BLOCK_SIZE_K * stride_bsk
c = accumulator.to(c_ptr.type.element_ty)
offs_cm = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M).to(tl.int64)
offs_cn = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N).to(tl.int64)
c_ptrs = (
c_ptr
+ stride_cm * offs_cm[:, None]
+ stride_cn * offs_cn[None, :]
+ pid_k * stride_ck
)
c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < N)
tl.store(c_ptrs, c, mask=c_mask)
_STATE = {}
_REDUCE_KERNEL = None
_MOD = None
_TRITON_KEY = (16, 2112, 7168)
def _ensure_imports():
global _MOD
if _MOD is not None:
return _MOD
import aiter
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle
_MOD = {
"dtypes": dtypes,
"quant": dynamic_mxfp4_quant,
"shuffle": e8m0_shuffle,
"gemm": aiter.gemm_a4w4,
}
return _MOD
def _get_reduce_kernel():
global _REDUCE_KERNEL
if _REDUCE_KERNEL is None:
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4 import (
_gemm_afp4wfp4_reduce_kernel,
)
_REDUCE_KERNEL = _gemm_afp4wfp4_reduce_kernel
return _REDUCE_KERNEL
def _reshape_b_shuffle(B_shuffle):
if not B_shuffle.is_contiguous():
B_shuffle = B_shuffle.contiguous()
return B_shuffle.view(torch.uint8).reshape(B_shuffle.shape[0] // 16, B_shuffle.shape[1] * 16)
def _reshape_b_scale(B_scale_sh):
if not B_scale_sh.is_contiguous():
B_scale_sh = B_scale_sh.contiguous()
return B_scale_sh.view(torch.uint8).reshape(B_scale_sh.shape[0] // 32, B_scale_sh.shape[1] * 32)
def _run_triton_kernel(A, B_shuffle, B_scale_sh, s):
c_target = s["y_pp"] if s["use_splitk"] else s["y"]
_fixed_preshuffle_kernel[s["grid"]](
A,
B_shuffle,
c_target,
B_scale_sh,
s["m"],
s["n"],
s["K"],
s["stride_am"],
s["stride_ak"],
s["stride_bn"],
s["stride_bk"],
s["stride_ck"],
s["stride_cm"],
s["stride_cn"],
s["stride_bsn"],
s["stride_bsk"],
**s["kernel_config"],
)
if s["use_splitk"]:
rk = _get_reduce_kernel()
rk[s["reduce_grid"]](
s["y_pp"],
s["y"],
s["m"],
s["n"],
s["y_pp"].stride(0),
s["y_pp"].stride(1),
s["y_pp"].stride(2),
s["stride_y_m"],
s["stride_y_n"],
16,
64,
s["actual_ksplit"],
s["max_ksplit"],
)
return s["y"]
def _init_shape2(A, B_shuffle, B_scale_sh):
from aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4 import _get_config
from aiter.ops.triton.gemm.basic.gemm_afp4wfp4 import get_splitk
m, k_bf16 = A.shape
n = B_shuffle.shape[0]
K = k_bf16 // 2
B_shuffle_u8 = _reshape_b_shuffle(B_shuffle)
B_scale_u8 = _reshape_b_scale(B_scale_sh)
config, _ = _get_config(m, n, K, True)
config["NUM_KSPLIT"] = 7
if config["NUM_KSPLIT"] > 1:
sbs, bk, nk = get_splitk(K, config["BLOCK_SIZE_K"], config["NUM_KSPLIT"])
config["SPLITK_BLOCK_SIZE"] = sbs
config["BLOCK_SIZE_K"] = bk
config["NUM_KSPLIT"] = nk
if config["BLOCK_SIZE_K"] >= 2 * K:
config["BLOCK_SIZE_K"] = triton.next_power_of_2(2 * K)
config["SPLITK_BLOCK_SIZE"] = 2 * K
config["NUM_KSPLIT"] = 1
config["BLOCK_SIZE_N"] = max(config["BLOCK_SIZE_N"], 32)
use_splitk = config["NUM_KSPLIT"] > 1
if not use_splitk:
config["SPLITK_BLOCK_SIZE"] = 2 * K
y = torch.empty((m, n), dtype=torch.bfloat16, device=A.device)
y_pp = None
if use_splitk:
y_pp = torch.empty((config["NUM_KSPLIT"], m, n), dtype=torch.float32, device=A.device)
state = {
"m": m,
"n": n,
"K": K,
"kernel_config": {
"BLOCK_SIZE_M": config["BLOCK_SIZE_M"],
"BLOCK_SIZE_N": config["BLOCK_SIZE_N"],
"BLOCK_SIZE_K": config["BLOCK_SIZE_K"],
"GROUP_SIZE_M": config["GROUP_SIZE_M"],
"NUM_KSPLIT": config["NUM_KSPLIT"],
"SPLITK_BLOCK_SIZE": config["SPLITK_BLOCK_SIZE"],
"num_warps": config["num_warps"],
"num_stages": config["num_stages"],
"waves_per_eu": config["waves_per_eu"],
"matrix_instr_nonkdim": config["matrix_instr_nonkdim"],
"cache_modifier": config.get("cache_modifier"),
},
"y": y,
"y_pp": y_pp,
"use_splitk": use_splitk,
"grid": (
config["NUM_KSPLIT"]
* triton.cdiv(m, config["BLOCK_SIZE_M"])
* triton.cdiv(n, config["BLOCK_SIZE_N"]),
),
"stride_am": k_bf16,
"stride_ak": 1,
"stride_bn": B_shuffle_u8.stride(0),
"stride_bk": B_shuffle_u8.stride(1),
"stride_ck": 0 if not use_splitk else y_pp.stride(0),
"stride_cm": y.stride(0) if not use_splitk else y_pp.stride(1),
"stride_cn": y.stride(1) if not use_splitk else y_pp.stride(2),
"stride_bsn": B_scale_u8.stride(0),
"stride_bsk": B_scale_u8.stride(1),
"stride_y_m": y.stride(0),
"stride_y_n": y.stride(1),
"reduce_grid": None,
"actual_ksplit": None,
"max_ksplit": None,
}
if use_splitk:
state["actual_ksplit"] = triton.cdiv(K, config["SPLITK_BLOCK_SIZE"] // 2)
state["max_ksplit"] = triton.next_power_of_2(config["NUM_KSPLIT"])
state["reduce_grid"] = (triton.cdiv(m, 16), triton.cdiv(n, 64))
_STATE[_TRITON_KEY] = state
def _fallback_wrapper(data):
mod = _ensure_imports()
A, B, B_q, B_shuffle, B_scale_sh = data
if not A.is_contiguous():
A = A.contiguous()
A_q_u8, A_scale_u8 = mod["quant"](A)
A_scale_sh = mod["shuffle"](A_scale_u8)
return mod["gemm"](
A_q_u8.view(mod["dtypes"].fp4x2),
B_shuffle,
A_scale_sh.view(mod["dtypes"].fp8_e8m0),
B_scale_sh,
dtype=mod["dtypes"].bf16,
bpreshuffle=True,
)
def custom_kernel(data: input_t) -> output_t:
A, B, B_q, B_shuffle, B_scale_sh = data
A = A.contiguous()
key = (A.shape[0], B_shuffle.shape[0], A.shape[1])
if key != _TRITON_KEY:
return _fallback_wrapper(data)
if key not in _STATE:
_init_shape2(A, B_shuffle, B_scale_sh)
return _run_triton_kernel(A, _reshape_b_shuffle(B_shuffle), _reshape_b_scale(B_scale_sh), _STATE[key])
scrolls · 351 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