Skip to content
KernelIndex
Search⌘K

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
AMD MXFP4 GEMMsuite of 6 cases
AMD Instinct MI355X
20.2µs
#722 of 1143
2026-04-06

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-kAttempt 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