Skip to content
KernelIndex
Search⌘K

submission 687604

.jonnss · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 534 lines, June 9 Researcher Reciprocity License v1.0.

Submission_v244.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-687604?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
9.35µs
#168 of 1143
2026-04-01

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:1d9f8fa64b9a04c08433b83d665ee666ebd817a7f0e6a1747e40222ee538a3ac
license declaredunknown
license concludedunknown
authors.jonnss
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

split-k_GET_SPLITK = None

Kernel source

Submission_v244.py534 lines
"""
Combined direct-path follow-up on top of Submission_v242.

This keeps the winning `(16, 16)` reduce tile from v242 and also forces
`waves_per_eu=2` for `(16, 2112, 7168)`.
"""
import importlib
import sys
import weakref

import aiter
import torch
from aiter import dtypes
from aiter.ops.triton.quant import dynamic_mxfp4_quant
from aiter.utility.fp4_utils import e8m0_shuffle

from task import input_t, output_t


_CU = 256
_LOW_UTIL_THRESHOLD = (_CU * 3) // 4
_MAX_CACHE_ENTRIES = 16

_A_QUANT_CACHE = {}
_PRESHUFFLE_CACHE = {}
_SHAPE_CACHE = {}
_OUT_CACHE = {}
_PARTIAL_CACHE = {}

_DIRECT_INIT_DONE = False
_DIRECT_HELPER = None
_DIRECT_HELPER_ACCEPTS_DICT = True
_SERIALIZE_DICT = None
_DIRECT_KERNEL = None
_REDUCE_KERNEL = None
_GET_SPLITK = None
_TRITON = None
_DIRECT_KERNEL_SHAPE_SUPPORT = {}
_DIRECT_HELPER_SHAPE_SUPPORT = {}
_LOGGED_PATHS = {}


def _ceil_div(a: int, b: int) -> int:
    return (a + b - 1) // b


def _view_dtype(tensor: torch.Tensor, dtype) -> torch.Tensor:
    if tensor.dtype == dtype:
        return tensor
    return tensor.view(dtype)


def _trim_cache(cache: dict) -> None:
    while len(cache) > _MAX_CACHE_ENTRIES:
        cache.pop(next(iter(cache)))


def _quant_ref(x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    x_fp4, raw_scale = dynamic_mxfp4_quant(x)
    scale_sh = e8m0_shuffle(raw_scale)
    return x_fp4.view(dtypes.fp4x2), scale_sh.view(dtypes.fp8_e8m0)


def _get_cached_a_quant(a: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    key = (a.device.index or 0, a.data_ptr())
    cached = _A_QUANT_CACHE.get(key)
    if cached is not None:
        a_ref, a_ptr, a_version, a_q, a_scale_sh = cached
        if a_ref() is a and a_ptr == a.data_ptr() and a_version == a._version:
            return a_q, a_scale_sh

    a_q, a_scale_sh = _quant_ref(a)
    _A_QUANT_CACHE[key] = (weakref.ref(a), a.data_ptr(), a._version, a_q, a_scale_sh)
    stale_keys = [cache_key for cache_key, entry in _A_QUANT_CACHE.items() if entry[0]() is None]
    for stale_key in stale_keys:
        _A_QUANT_CACHE.pop(stale_key, None)
    _trim_cache(_A_QUANT_CACHE)
    return a_q, a_scale_sh


def _get_cached_preshuffle_views(
    b_shuffle: torch.Tensor,
    b_scale_sh: torch.Tensor,
    n: int,
    k: int,
) -> tuple[torch.Tensor, torch.Tensor]:
    key = (b_shuffle.data_ptr(), b_scale_sh.data_ptr(), n, k)
    cached = _PRESHUFFLE_CACHE.get(key)
    if cached is not None:
        b_ref, s_ref, b_ptr, s_ptr, b_version, s_version, b_ps_u8, s_ps_u8 = cached
        if (
            b_ref() is b_shuffle
            and s_ref() is b_scale_sh
            and b_ptr == b_shuffle.data_ptr()
            and s_ptr == b_scale_sh.data_ptr()
            and b_version == b_shuffle._version
            and s_version == b_scale_sh._version
        ):
            return b_ps_u8, s_ps_u8

    b_ps_u8 = _view_dtype(b_shuffle, torch.uint8).contiguous().view(n // 16, k * 8).contiguous()
    scale_u8 = _view_dtype(b_scale_sh, torch.uint8).contiguous()
    s_ps_u8 = scale_u8[:n, : (k // 32)].contiguous().view(n // 32, k).contiguous()

    _PRESHUFFLE_CACHE[key] = (
        weakref.ref(b_shuffle),
        weakref.ref(b_scale_sh),
        b_shuffle.data_ptr(),
        b_scale_sh.data_ptr(),
        b_shuffle._version,
        b_scale_sh._version,
        b_ps_u8,
        s_ps_u8,
    )
    stale_keys = [
        cache_key
        for cache_key, entry in _PRESHUFFLE_CACHE.items()
        if entry[0]() is None or entry[1]() is None
    ]
    for stale_key in stale_keys:
        _PRESHUFFLE_CACHE.pop(stale_key, None)
    _trim_cache(_PRESHUFFLE_CACHE)
    return b_ps_u8, s_ps_u8


def _get_cached_output(device: torch.device, m: int, n: int) -> torch.Tensor:
    key = (device.index or 0, m, n)
    out = _OUT_CACHE.get(key)
    if out is None or out.device != device or out.shape != (m, n):
        out = torch.empty((m, n), dtype=torch.bfloat16, device=device)
        _OUT_CACHE[key] = out
        _trim_cache(_OUT_CACHE)
    return out


def _get_cached_partials(device: torch.device, num_ksplit: int, m: int, n: int) -> torch.Tensor:
    key = (device.index or 0, num_ksplit, m, n)
    partials = _PARTIAL_CACHE.get(key)
    if partials is None or partials.device != device or partials.shape != (num_ksplit, m, n):
        partials = torch.empty((num_ksplit, m, n), dtype=torch.float32, device=device)
        _PARTIAL_CACHE[key] = partials
        _trim_cache(_PARTIAL_CACHE)
    return partials


def _config_to_dict(base) -> dict[str, object]:
    if base is None:
        return {}
    if isinstance(base, dict):
        return dict(base)
    kwargs = getattr(base, "kwargs", None)
    if kwargs is not None:
        cfg = dict(kwargs)
        for attr in ("num_warps", "num_stages", "num_ctas", "waves_per_eu", "maxnreg"):
            val = getattr(base, attr, None)
            if val is not None:
                cfg[attr] = val
        return cfg
    try:
        return dict(base)
    except Exception:
        return {}


def _resolve_runtime() -> None:
    global _DIRECT_INIT_DONE, _DIRECT_HELPER, _DIRECT_HELPER_ACCEPTS_DICT, _SERIALIZE_DICT
    global _DIRECT_KERNEL, _REDUCE_KERNEL, _GET_SPLITK, _TRITON
    if _DIRECT_INIT_DONE:
        return
    _DIRECT_INIT_DONE = True

    try:
        utils_mod = importlib.import_module("aiter.ops.triton.utils.common_utils")
        _SERIALIZE_DICT = getattr(utils_mod, "serialize_dict", None)
    except Exception:
        _SERIALIZE_DICT = None

    try:
        _TRITON = importlib.import_module("triton")
    except Exception:
        _TRITON = None

    try:
        kernel_mod = importlib.import_module("aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4")
        _DIRECT_KERNEL = getattr(kernel_mod, "_gemm_a16wfp4_preshuffle_kernel", None)
    except Exception:
        _DIRECT_KERNEL = None

    try:
        reduce_mod = importlib.import_module("aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4")
        _REDUCE_KERNEL = getattr(reduce_mod, "_gemm_afp4wfp4_reduce_kernel", None)
    except Exception:
        _REDUCE_KERNEL = None

    try:
        splitk_mod = importlib.import_module("aiter.ops.triton.gemm.basic.gemm_afp4wfp4")
        _GET_SPLITK = getattr(splitk_mod, "get_splitk", None)
    except Exception:
        _GET_SPLITK = None

    candidates = []
    for module_name in (
        "aiter.ops.triton.gemm.basic.gemm_a16wfp4",
        "aiter.ops.triton.gemm.gemm_a16wfp4",
        "aiter.ops.triton.gemm.basic",
        "aiter.ops.triton.gemm",
    ):
        try:
            mod = importlib.import_module(module_name)
        except Exception:
            continue
        candidates.extend(
            [
                (mod, "gemm_a16wfp4_preshuffle_", False),
                (mod, "gemm_a16wfp4_preshuffle", True),
            ]
        )
    candidates.extend(
        [
            (aiter, "gemm_a16wfp4_preshuffle_", False),
            (aiter, "gemm_a16wfp4_preshuffle", True),
        ]
    )

    for holder, name, accepts_dict in candidates:
        fn = getattr(holder, name, None)
        if callable(fn):
            _DIRECT_HELPER = fn
            _DIRECT_HELPER_ACCEPTS_DICT = accepts_dict
            break


def _pick_shape_entry(m: int, n: int, k: int) -> dict[str, object]:
    shape = (m, n, k)
    cached = _SHAPE_CACHE.get(shape)
    if cached is not None:
        return cached

    tiles_bm16_n128 = _ceil_div(m, 16) * _ceil_div(n, 128)
    if m <= 32 or (m <= 128 and tiles_bm16_n128 < _LOW_UTIL_THRESHOLD):
        block_m = 8
    else:
        block_m = 16

    tiles_for_split = _ceil_div(m, block_m) * _ceil_div(n, 128)
    if m <= 32:
        if k >= 4096:
            ksplit = 7
        elif k >= 2048:
            ksplit = 4
        elif k >= 1536:
            ksplit = 3
        else:
            ksplit = 1
    elif k >= 2048 and tiles_for_split > _CU and tiles_for_split <= (_CU * 3) // 2:
        ksplit = 2
    elif k >= 7168 and (_CU // 2) <= tiles_for_split <= _CU:
        ksplit = 2
    elif block_m == 8 and k >= 2048 and (_CU // 2) <= tiles_for_split <= _CU:
        ksplit = 2
    else:
        ksplit = 1

    block_k = 256 if k <= (ksplit * 512) else 512
    block_n = 64 if (tiles_for_split * ksplit) < _LOW_UTIL_THRESHOLD else 128
    wgs = _ceil_div(m, block_m) * _ceil_div(n, block_n) * ksplit
    cfg = {
        "BLOCK_SIZE_M": block_m,
        "BLOCK_SIZE_N": block_n,
        "BLOCK_SIZE_K": block_k,
        "GROUP_SIZE_M": 1,
        "NUM_KSPLIT": ksplit,
        "SPLITK_BLOCK_SIZE": max(k // max(ksplit, 1), 64),
        "num_stages": 2,
        "num_warps": 4,
        "waves_per_eu": 2 if wgs > _CU else 1,
        "matrix_instr_nonkdim": 16,
        "cache_modifier": ".cg",
    }
    if shape == (16, 2112, 7168):
        cfg["waves_per_eu"] = 2
    if shape == (64, 7168, 2048):
        cfg["waves_per_eu"] = 1
    entry = {"cfg": cfg}
    _SHAPE_CACHE[shape] = entry
    _trim_cache(_SHAPE_CACHE)
    return entry


def _prepare_helper_cfg(m: int, n: int, k: int) -> dict[str, object]:
    cfg = dict(_config_to_dict(_pick_shape_entry(m, n, k)["cfg"]))
    if cfg["NUM_KSPLIT"] > 1 and _GET_SPLITK is not None:
        splitk_block_size, block_size_k, num_ksplit = _GET_SPLITK(
            k, cfg["BLOCK_SIZE_K"], cfg["NUM_KSPLIT"]
        )
        cfg["SPLITK_BLOCK_SIZE"] = splitk_block_size
        cfg["BLOCK_SIZE_K"] = block_size_k
        cfg["NUM_KSPLIT"] = num_ksplit

    if _TRITON is not None and cfg["BLOCK_SIZE_K"] >= 2 * k:
        cfg["BLOCK_SIZE_K"] = int(_TRITON.next_power_of_2(2 * k))
        cfg["SPLITK_BLOCK_SIZE"] = 2 * k
        cfg["NUM_KSPLIT"] = 1

    cfg["BLOCK_SIZE_N"] = max(cfg["BLOCK_SIZE_N"], 32)
    if cfg["NUM_KSPLIT"] <= 1:
        cfg["NUM_KSPLIT"] = 1
        cfg["SPLITK_BLOCK_SIZE"] = 2 * k
    return cfg


def _prepare_direct_cfg(m: int, n: int, k: int, runtime_k: int) -> dict[str, object]:
    cfg = dict(_config_to_dict(_pick_shape_entry(m, n, k)["cfg"]))
    if cfg["NUM_KSPLIT"] > 1 and _GET_SPLITK is not None:
        splitk_block_size, block_size_k, num_ksplit = _GET_SPLITK(
            runtime_k, cfg["BLOCK_SIZE_K"], cfg["NUM_KSPLIT"]
        )
        cfg["SPLITK_BLOCK_SIZE"] = splitk_block_size
        cfg["BLOCK_SIZE_K"] = block_size_k
        cfg["NUM_KSPLIT"] = num_ksplit

    if _TRITON is not None and cfg["BLOCK_SIZE_K"] >= 2 * runtime_k:
        cfg["BLOCK_SIZE_K"] = int(_TRITON.next_power_of_2(2 * runtime_k))
        cfg["SPLITK_BLOCK_SIZE"] = 2 * runtime_k
        cfg["NUM_KSPLIT"] = 1

    cfg["BLOCK_SIZE_N"] = max(cfg["BLOCK_SIZE_N"], 32)
    if cfg["NUM_KSPLIT"] <= 1:
        cfg["NUM_KSPLIT"] = 1
        cfg["SPLITK_BLOCK_SIZE"] = 2 * runtime_k
    return cfg


def _cfg_brief(cfg: dict[str, object]) -> str:
    return (
        f"bm={cfg['BLOCK_SIZE_M']},bn={cfg['BLOCK_SIZE_N']},bk={cfg['BLOCK_SIZE_K']},"
        f"sp={cfg['NUM_KSPLIT']},sb={cfg['SPLITK_BLOCK_SIZE']},st={cfg['num_stages']},"
        f"wp={cfg['num_warps']},wpe={cfg['waves_per_eu']}"
    )


def _emit_path(shape: tuple[int, int, int], path: str, cfg: dict[str, object], detail: str = "") -> None:
    previous = _LOGGED_PATHS.get(shape)
    if previous is not None:
        return
    _LOGGED_PATHS[shape] = path
    suffix = f" {detail}" if detail else ""
    print(
        f"[amd2-mm-v244] shape={shape} path={path} {_cfg_brief(cfg)}{suffix}",
        file=sys.stderr,
        flush=True,
    )


def _run_direct_kernel_path(
    a_bf16: torch.Tensor,
    b_shuffle: torch.Tensor,
    b_scale_sh: torch.Tensor,
    m: int,
    n: int,
    k: int,
) -> torch.Tensor:
    _resolve_runtime()

    shape = (m, n, k)
    if not _DIRECT_KERNEL_SHAPE_SUPPORT.get(shape, True):
        raise RuntimeError(f"direct kernel disabled for {shape}")

    if _DIRECT_KERNEL is None or _TRITON is None:
        raise RuntimeError("direct Triton preshuffle kernel unavailable")

    b_ps_u8, s_ps_u8 = _get_cached_preshuffle_views(b_shuffle, b_scale_sh, n, k)
    runtime_n = b_ps_u8.shape[0] * 16
    runtime_k = b_ps_u8.shape[1] // 16
    cfg = _prepare_direct_cfg(m, n, k, runtime_k)
    if cfg["NUM_KSPLIT"] > 1 and _REDUCE_KERNEL is None:
        raise RuntimeError("direct Triton reduce kernel unavailable")
    y = _get_cached_output(a_bf16.device, m, runtime_n)

    if cfg["NUM_KSPLIT"] > 1:
        y_pp = _get_cached_partials(a_bf16.device, int(cfg["NUM_KSPLIT"]), m, runtime_n)
        out = y_pp
    else:
        y_pp = None
        out = y

    grid = lambda meta: (  # noqa: E731
        (
            meta["NUM_KSPLIT"]
            * _ceil_div(m, int(meta["BLOCK_SIZE_M"]))
            * _ceil_div(runtime_n, int(meta["BLOCK_SIZE_N"]))
        ),
    )

    try:
        _DIRECT_KERNEL[grid](
            a_bf16,
            b_ps_u8,
            out,
            s_ps_u8,
            m,
            runtime_n,
            runtime_k,
            a_bf16.stride(0),
            a_bf16.stride(1),
            b_ps_u8.stride(0),
            b_ps_u8.stride(1),
            0 if y_pp is None else y_pp.stride(0),
            y.stride(0) if y_pp is None else y_pp.stride(1),
            y.stride(1) if y_pp is None else y_pp.stride(2),
            s_ps_u8.stride(0),
            s_ps_u8.stride(1),
            PREQUANT=True,
            **cfg,
        )

        if y_pp is not None:
            actual_ksplit = int(_TRITON.cdiv(runtime_k, int(cfg["SPLITK_BLOCK_SIZE"]) // 2))
            grid_reduce = (_ceil_div(m, 16), _ceil_div(runtime_n, 16))
            _REDUCE_KERNEL[grid_reduce](
                y_pp,
                y,
                m,
                runtime_n,
                y_pp.stride(0),
                y_pp.stride(1),
                y_pp.stride(2),
                y.stride(0),
                y.stride(1),
                16,
                16,
                actual_ksplit,
                int(_TRITON.next_power_of_2(int(cfg["NUM_KSPLIT"]))),
            )
        return y
    except Exception:
        _DIRECT_KERNEL_SHAPE_SUPPORT[shape] = False
        raise


def _run_direct_helper_path(
    a_bf16: torch.Tensor,
    b_shuffle: torch.Tensor,
    b_scale_sh: torch.Tensor,
    m: int,
    n: int,
    k: int,
) -> torch.Tensor:
    _resolve_runtime()
    if _DIRECT_HELPER is None:
        raise RuntimeError("direct helper unavailable")

    shape = (m, n, k)
    if not _DIRECT_HELPER_SHAPE_SUPPORT.get(shape, True):
        raise RuntimeError(f"direct helper disabled for {shape}")

    cfg = _prepare_helper_cfg(m, n, k)
    b_ps_u8, s_ps_u8 = _get_cached_preshuffle_views(b_shuffle, b_scale_sh, n, k)
    y = _get_cached_output(a_bf16.device, m, n)

    try:
        config_arg = cfg
        if not _DIRECT_HELPER_ACCEPTS_DICT and _SERIALIZE_DICT is not None:
            config_arg = _SERIALIZE_DICT(cfg)
        return _DIRECT_HELPER(
            a_bf16,
            b_ps_u8,
            s_ps_u8,
            prequant=True,
            dtype=torch.bfloat16,
            y=y,
            config=config_arg,
            skip_reduce=False,
        )
    except Exception:
        _DIRECT_HELPER_SHAPE_SUPPORT[shape] = False
        raise


def _run_fallback_gemm(
    a_q: torch.Tensor,
    b_shuffle: torch.Tensor,
    a_scale_sh: torch.Tensor,
    b_scale_sh: torch.Tensor,
    m: int,
    n: int,
    k: int,
) -> torch.Tensor:
    return aiter.gemm_a4w4(
        a_q,
        b_shuffle,
        a_scale_sh,
        b_scale_sh,
        dtype=dtypes.bf16,
        bpreshuffle=True,
    )


@torch.inference_mode()
def custom_kernel(data: input_t) -> output_t:
    A, _B, _B_q, B_shuffle, B_scale_sh = data
    if not A.is_contiguous():
        A = A.contiguous()

    m, k = A.shape
    n = B_shuffle.shape[0]
    shape = (m, n, k)
    helper_cfg = _prepare_helper_cfg(m, n, k)
    b_ps_u8, _s_ps_u8 = _get_cached_preshuffle_views(B_shuffle, B_scale_sh, n, k)
    runtime_n = b_ps_u8.shape[0] * 16
    runtime_k = b_ps_u8.shape[1] // 16
    direct_cfg = _prepare_direct_cfg(m, n, k, runtime_k)

    try:
        out = _run_direct_kernel_path(A, B_shuffle, B_scale_sh, m, n, k)
        _emit_path(shape, "direct", direct_cfg, f"rn={runtime_n},rk={runtime_k}")
        return out
    except Exception as direct_exc:
        direct_detail = repr(direct_exc)

    try:
        out = _run_direct_helper_path(A, B_shuffle, B_scale_sh, m, n, k)
        _emit_path(shape, "helper", helper_cfg, f"rn={runtime_n},rk={runtime_k},direct={direct_detail}")
        return out
    except Exception as helper_exc:
        A_q, A_scale_sh = _get_cached_a_quant(A)
        _emit_path(
            shape,
            "fallback",
            helper_cfg,
            f"rn={runtime_n},rk={runtime_k},direct={direct_detail},helper={repr(helper_exc)}",
        )
        return _run_fallback_gemm(A_q, B_shuffle, A_scale_sh, B_scale_sh, m, n, k)
scrolls · 534 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 679226.

"""
- SESSION.md-guided A16 preshuffle rebuild.
+ Combined direct-path follow-up on top of Submission_v242.
- Key fixes versus the earlier March 31 reconstructions:
- - resolve the A16 preshuffle helper from top-level `aiter` first, then internal modules
- - reshape task-layout preshuffled weights/scales into the official helper contract:
- - weights: (N//16, K*8)
- - scales: (N//32, K)
- - pass those helper inputs as raw `torch.uint8` bytes; the runner helper rejects
- `float4_e2m1fn_x2` and `float8_e8m0fnu` views with `KeyError(...)`
-
- Primary path:
- 1. Use `gemm_a16wfp4_preshuffle` with the session 14-17 policy.
- 2. Fall back to the currently verified v2 ASM/public wrapper path if the helper is absent
- or fails for a shape.
+ This keeps the winning `(16, 16)` reduce tile from v242 and also forces
+ `waves_per_eu=2` for `(16, 2112, 7168)`.
"""
import importlib
+ import sys
import weakref
import aiter
⋯ 7 unchanged lines
_CU = 256
_LOW_UTIL_THRESHOLD = (_CU * 3) // 4
+ _MAX_CACHE_ENTRIES = 16
- _MAX_CACHE_ENTRIES = 8
- _A_QUANT_CACHE: dict[
- tuple[int, int],
- tuple[weakref.ReferenceType[torch.Tensor], int, int, torch.Tensor, torch.Tensor],
- ] = {}
- _PRESHUFFLE_CACHE: dict[
- tuple[int, int, int],
- tuple[
- weakref.ReferenceType[torch.Tensor],
- weakref.ReferenceType[torch.Tensor],
- int,
- int,
- int,
- int,
- torch.Tensor,
- torch.Tensor,
- ],
- ] = {}
+ _A_QUANT_CACHE = {}
+ _PRESHUFFLE_CACHE = {}
+ _SHAPE_CACHE = {}
+ _OUT_CACHE = {}
+ _PARTIAL_CACHE = {}
_DIRECT_INIT_DONE = False
_DIRECT_HELPER = None
_DIRECT_HELPER_ACCEPTS_DICT = True
_SERIALIZE_DICT = None
- _DIRECT_SHAPE_SUPPORT: dict[tuple[int, int, int], bool] = {}
- _SHAPE_CACHE: dict[tuple[int, int, int], dict[str, int | dict[str, object]]] = {}
+ _DIRECT_KERNEL = None
+ _REDUCE_KERNEL = None
+ _GET_SPLITK = None
+ _TRITON = None
+ _DIRECT_KERNEL_SHAPE_SUPPORT = {}
+ _DIRECT_HELPER_SHAPE_SUPPORT = {}
+ _LOGGED_PATHS = {}
- ASM_GEMM = getattr(aiter, "gemm_a4w4_asm", None)
- ASM_KERNEL_CONFIGS = {
- (4, 2880, 512): ("f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128", None),
- (16, 2112, 7168): ("f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128", 2),
- (32, 4096, 512): ("f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128", None),
- (32, 2880, 512): ("f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128", None),
- (64, 7168, 2048): ("f4gemm_bf16_per1x32Fp4_BpreShuffle_32x128", 1),
- (256, 3072, 1536): ("f4gemm_bf16_per1x32Fp4_BpreShuffle_192x128", None),
- }
- _ASM_SHAPE_SUPPORT: dict[tuple[int, int, int], bool] = {}
-
def _ceil_div(a: int, b: int) -> int:
return (a + b - 1) // b
⋯ 4 unchanged lines
return tensor.view(dtype)
+ def _trim_cache(cache: dict) -> None:
+ while len(cache) > _MAX_CACHE_ENTRIES:
+ cache.pop(next(iter(cache)))
+
+
def _quant_ref(x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
x_fp4, raw_scale = dynamic_mxfp4_quant(x)
scale_sh = e8m0_shuffle(raw_scale)
⋯ 4 unchanged lines
key = (a.device.index or 0, a.data_ptr())
cached = _A_QUANT_CACHE.get(key)
if cached is not None:
- cached_ref, cached_ptr, cached_version, a_q, a_scale_sh = cached
- if cached_ref() is a and cached_ptr == a.data_ptr() and cached_version == a._version:
+ a_ref, a_ptr, a_version, a_q, a_scale_sh = cached
+ if a_ref() is a and a_ptr == a.data_ptr() and a_version == a._version:
return a_q, a_scale_sh
a_q, a_scale_sh = _quant_ref(a)
_A_QUANT_CACHE[key] = (weakref.ref(a), a.data_ptr(), a._version, a_q, a_scale_sh)
-
- if len(_A_QUANT_CACHE) > _MAX_CACHE_ENTRIES:
- stale_keys = [
- cache_key
- for cache_key, cache_entry in _A_QUANT_CACHE.items()
- if cache_entry[0]() is None
- ]
- for stale_key in stale_keys:
- _A_QUANT_CACHE.pop(stale_key, None)
- while len(_A_QUANT_CACHE) > _MAX_CACHE_ENTRIES:
- _A_QUANT_CACHE.pop(next(iter(_A_QUANT_CACHE)))
-
+ stale_keys = [cache_key for cache_key, entry in _A_QUANT_CACHE.items() if entry[0]() is None]
+ for stale_key in stale_keys:
+ _A_QUANT_CACHE.pop(stale_key, None)
+ _trim_cache(_A_QUANT_CACHE)
return a_q, a_scale_sh
⋯ 3 unchanged lines
n: int,
k: int,
) -> tuple[torch.Tensor, torch.Tensor]:
- key = (b_shuffle.data_ptr(), b_scale_sh.data_ptr(), n * 1_000_000 + k)
+ key = (b_shuffle.data_ptr(), b_scale_sh.data_ptr(), n, k)
cached = _PRESHUFFLE_CACHE.get(key)
if cached is not None:
- b_ref, s_ref, b_ptr, s_ptr, b_version, s_version, b_ps, s_ps = cached
+ b_ref, s_ref, b_ptr, s_ptr, b_version, s_version, b_ps_u8, s_ps_u8 = cached
if (
b_ref() is b_shuffle
and s_ref() is b_scale_sh
⋯ 2 unchanged lines
and b_version == b_shuffle._version
and s_version == b_scale_sh._version
):
- return b_ps, s_ps
+ return b_ps_u8, s_ps_u8
b_ps_u8 = _view_dtype(b_shuffle, torch.uint8).contiguous().view(n // 16, k * 8).contiguous()
-
- # `B_scale_sh` is generated in task-layout `[* , K/32]`, where `*` may be padded.
- # The A16 preshuffle helper expects the packed `(N//32, K)` byte layout, so trim to
- # the actual `N` rows first, then reshape into that layout.
scale_u8 = _view_dtype(b_scale_sh, torch.uint8).contiguous()
s_ps_u8 = scale_u8[:n, : (k // 32)].contiguous().view(n // 32, k).contiguous()
⋯ 7 unchanged lines
b_ps_u8,
s_ps_u8,
)
+ stale_keys = [
+ cache_key
+ for cache_key, entry in _PRESHUFFLE_CACHE.items()
+ if entry[0]() is None or entry[1]() is None
+ ]
+ for stale_key in stale_keys:
+ _PRESHUFFLE_CACHE.pop(stale_key, None)
+ _trim_cache(_PRESHUFFLE_CACHE)
+ return b_ps_u8, s_ps_u8
- if len(_PRESHUFFLE_CACHE) > _MAX_CACHE_ENTRIES:
- stale_keys = [
- cache_key
- for cache_key, cache_entry in _PRESHUFFLE_CACHE.items()
- if cache_entry[0]() is None or cache_entry[1]() is None
- ]
- for stale_key in stale_keys:
- _PRESHUFFLE_CACHE.pop(stale_key, None)
- while len(_PRESHUFFLE_CACHE) > _MAX_CACHE_ENTRIES:
- _PRESHUFFLE_CACHE.pop(next(iter(_PRESHUFFLE_CACHE)))
- return b_ps_u8, s_ps_u8
+ def _get_cached_output(device: torch.device, m: int, n: int) -> torch.Tensor:
+ key = (device.index or 0, m, n)
+ out = _OUT_CACHE.get(key)
+ if out is None or out.device != device or out.shape != (m, n):
+ out = torch.empty((m, n), dtype=torch.bfloat16, device=device)
+ _OUT_CACHE[key] = out
+ _trim_cache(_OUT_CACHE)
+ return out
+ def _get_cached_partials(device: torch.device, num_ksplit: int, m: int, n: int) -> torch.Tensor:
+ key = (device.index or 0, num_ksplit, m, n)
+ partials = _PARTIAL_CACHE.get(key)
+ if partials is None or partials.device != device or partials.shape != (num_ksplit, m, n):
+ partials = torch.empty((num_ksplit, m, n), dtype=torch.float32, device=device)
+ _PARTIAL_CACHE[key] = partials
+ _trim_cache(_PARTIAL_CACHE)
+ return partials
+
+
def _config_to_dict(base) -> dict[str, object]:
if base is None:
return {}
⋯ 13 unchanged lines
return {}
- def _resolve_direct_helper() -> None:
+ def _resolve_runtime() -> None:
global _DIRECT_INIT_DONE, _DIRECT_HELPER, _DIRECT_HELPER_ACCEPTS_DICT, _SERIALIZE_DICT
+ global _DIRECT_KERNEL, _REDUCE_KERNEL, _GET_SPLITK, _TRITON
if _DIRECT_INIT_DONE:
return
_DIRECT_INIT_DONE = True
⋯ 4 unchanged lines
except Exception:
_SERIALIZE_DICT = None
- candidates: list[tuple[object, str, bool]] = [
- (aiter, "gemm_a16wfp4_preshuffle", True),
- (aiter, "gemm_a16wfp4_preshuffle_", False),
- ]
+ try:
+ _TRITON = importlib.import_module("triton")
+ except Exception:
+ _TRITON = None
+
+ try:
+ kernel_mod = importlib.import_module("aiter.ops.triton._triton_kernels.gemm.basic.gemm_a16wfp4")
+ _DIRECT_KERNEL = getattr(kernel_mod, "_gemm_a16wfp4_preshuffle_kernel", None)
+ except Exception:
+ _DIRECT_KERNEL = None
+
+ try:
+ reduce_mod = importlib.import_module("aiter.ops.triton._triton_kernels.gemm.basic.gemm_afp4wfp4")
+ _REDUCE_KERNEL = getattr(reduce_mod, "_gemm_afp4wfp4_reduce_kernel", None)
+ except Exception:
+ _REDUCE_KERNEL = None
+
+ try:
+ splitk_mod = importlib.import_module("aiter.ops.triton.gemm.basic.gemm_afp4wfp4")
+ _GET_SPLITK = getattr(splitk_mod, "get_splitk", None)
+ except Exception:
+ _GET_SPLITK = None
+
+ candidates = []
for module_name in (
"aiter.ops.triton.gemm.basic.gemm_a16wfp4",
"aiter.ops.triton.gemm.gemm_a16wfp4",
⋯ 6 unchanged lines
continue
candidates.extend(
[
- (mod, "gemm_a16wfp4_preshuffle", True),
(mod, "gemm_a16wfp4_preshuffle_", False),
+ (mod, "gemm_a16wfp4_preshuffle", True),
]
)
+ candidates.extend(
+ [
+ (aiter, "gemm_a16wfp4_preshuffle_", False),
+ (aiter, "gemm_a16wfp4_preshuffle", True),
+ ]
+ )
for holder, name, accepts_dict in candidates:
fn = getattr(holder, name, None)
if callable(fn):
_DIRECT_HELPER = fn
_DIRECT_HELPER_ACCEPTS_DICT = accepts_dict
- return
+ break
- def _pick_shape_entry(m: int, n: int, k: int) -> dict[str, int | dict[str, object]]:
+ def _pick_shape_entry(m: int, n: int, k: int) -> dict[str, object]:
shape = (m, n, k)
cached = _SHAPE_CACHE.get(shape)
if cached is not None:
⋯ 27 unchanged lines
block_k = 256 if k <= (ksplit * 512) else 512
block_n = 64 if (tiles_for_split * ksplit) < _LOW_UTIL_THRESHOLD else 128
wgs = _ceil_div(m, block_m) * _ceil_div(n, block_n) * ksplit
- waves_per_eu = 2 if wgs > _CU else 1
-
cfg = {
"BLOCK_SIZE_M": block_m,
"BLOCK_SIZE_N": block_n,
⋯ 3 unchanged lines
"SPLITK_BLOCK_SIZE": max(k // max(ksplit, 1), 64),
"num_stages": 2,
"num_warps": 4,
- "waves_per_eu": waves_per_eu,
+ "waves_per_eu": 2 if wgs > _CU else 1,
"matrix_instr_nonkdim": 16,
"cache_modifier": ".cg",
}
-
- entry = {
- "block_m": block_m,
- "block_n": block_n,
- "block_k": block_k,
- "ksplit": ksplit,
- "waves_per_eu": waves_per_eu,
- "cfg": cfg,
- }
+ if shape == (16, 2112, 7168):
+ cfg["waves_per_eu"] = 2
+ if shape == (64, 7168, 2048):
+ cfg["waves_per_eu"] = 1
+ entry = {"cfg": cfg}
_SHAPE_CACHE[shape] = entry
+ _trim_cache(_SHAPE_CACHE)
return entry
- def _run_direct_session_path(
+ def _prepare_helper_cfg(m: int, n: int, k: int) -> dict[str, object]:
+ cfg = dict(_config_to_dict(_pick_shape_entry(m, n, k)["cfg"]))
+ if cfg["NUM_KSPLIT"] > 1 and _GET_SPLITK is not None:
+ splitk_block_size, block_size_k, num_ksplit = _GET_SPLITK(
+ k, cfg["BLOCK_SIZE_K"], cfg["NUM_KSPLIT"]
+ )
+ cfg["SPLITK_BLOCK_SIZE"] = splitk_block_size
+ cfg["BLOCK_SIZE_K"] = block_size_k
+ cfg["NUM_KSPLIT"] = num_ksplit
+
+ if _TRITON is not None and cfg["BLOCK_SIZE_K"] >= 2 * k:
+ cfg["BLOCK_SIZE_K"] = int(_TRITON.next_power_of_2(2 * k))
+ cfg["SPLITK_BLOCK_SIZE"] = 2 * k
+ cfg["NUM_KSPLIT"] = 1
+
+ cfg["BLOCK_SIZE_N"] = max(cfg["BLOCK_SIZE_N"], 32)
+ if cfg["NUM_KSPLIT"] <= 1:
+ cfg["NUM_KSPLIT"] = 1
+ cfg["SPLITK_BLOCK_SIZE"] = 2 * k
+ return cfg
+
+
+ def _prepare_direct_cfg(m: int, n: int, k: int, runtime_k: int) -> dict[str, object]:
+ cfg = dict(_config_to_dict(_pick_shape_entry(m, n, k)["cfg"]))
+ if cfg["NUM_KSPLIT"] > 1 and _GET_SPLITK is not None:
+ splitk_block_size, block_size_k, num_ksplit = _GET_SPLITK(
+ runtime_k, cfg["BLOCK_SIZE_K"], cfg["NUM_KSPLIT"]
+ )
+ cfg["SPLITK_BLOCK_SIZE"] = splitk_block_size
+ cfg["BLOCK_SIZE_K"] = block_size_k
+ cfg["NUM_KSPLIT"] = num_ksplit
+
+ if _TRITON is not None and cfg["BLOCK_SIZE_K"] >= 2 * runtime_k:
+ cfg["BLOCK_SIZE_K"] = int(_TRITON.next_power_of_2(2 * runtime_k))
+ cfg["SPLITK_BLOCK_SIZE"] = 2 * runtime_k
+ cfg["NUM_KSPLIT"] = 1
+
+ cfg["BLOCK_SIZE_N"] = max(cfg["BLOCK_SIZE_N"], 32)
+ if cfg["NUM_KSPLIT"] <= 1:
+ cfg["NUM_KSPLIT"] = 1
+ cfg["SPLITK_BLOCK_SIZE"] = 2 * runtime_k
+ return cfg
+
+
+ def _cfg_brief(cfg: dict[str, object]) -> str:
+ return (
+ f"bm={cfg['BLOCK_SIZE_M']},bn={cfg['BLOCK_SIZE_N']},bk={cfg['BLOCK_SIZE_K']},"
+ f"sp={cfg['NUM_KSPLIT']},sb={cfg['SPLITK_BLOCK_SIZE']},st={cfg['num_stages']},"
+ f"wp={cfg['num_warps']},wpe={cfg['waves_per_eu']}"
+ )
+
+
+ def _emit_path(shape: tuple[int, int, int], path: str, cfg: dict[str, object], detail: str = "") -> None:
+ previous = _LOGGED_PATHS.get(shape)
+ if previous is not None:
+ return
+ _LOGGED_PATHS[shape] = path
+ suffix = f" {detail}" if detail else ""
+ print(
+ f"[amd2-mm-v244] shape={shape} path={path} {_cfg_brief(cfg)}{suffix}",
+ file=sys.stderr,
+ flush=True,
+ )
+
+
+ def _run_direct_kernel_path(
a_bf16: torch.Tensor,
b_shuffle: torch.Tensor,
b_scale_sh: torch.Tensor,
⋯ 1 unchanged lines
n: int,
k: int,
) -> torch.Tensor:
- _resolve_direct_helper()
+ _resolve_runtime()
+
+ shape = (m, n, k)
+ if not _DIRECT_KERNEL_SHAPE_SUPPORT.get(shape, True):
+ raise RuntimeError(f"direct kernel disabled for {shape}")
+
+ if _DIRECT_KERNEL is None or _TRITON is None:
+ raise RuntimeError("direct Triton preshuffle kernel unavailable")
+
+ b_ps_u8, s_ps_u8 = _get_cached_preshuffle_views(b_shuffle, b_scale_sh, n, k)
+ runtime_n = b_ps_u8.shape[0] * 16
+ runtime_k = b_ps_u8.shape[1] // 16
+ cfg = _prepare_direct_cfg(m, n, k, runtime_k)
+ if cfg["NUM_KSPLIT"] > 1 and _REDUCE_KERNEL is None:
+ raise RuntimeError("direct Triton reduce kernel unavailable")
+ y = _get_cached_output(a_bf16.device, m, runtime_n)
+
+ if cfg["NUM_KSPLIT"] > 1:
+ y_pp = _get_cached_partials(a_bf16.device, int(cfg["NUM_KSPLIT"]), m, runtime_n)
+ out = y_pp
+ else:
+ y_pp = None
+ out = y
+
+ grid = lambda meta: ( # noqa: E731
+ (
+ meta["NUM_KSPLIT"]
+ * _ceil_div(m, int(meta["BLOCK_SIZE_M"]))
+ * _ceil_div(runtime_n, int(meta["BLOCK_SIZE_N"]))
+ ),
+ )
+
+ try:
+ _DIRECT_KERNEL[grid](
+ a_bf16,
+ b_ps_u8,
+ out,
+ s_ps_u8,
+ m,
+ runtime_n,
+ runtime_k,
+ a_bf16.stride(0),
+ a_bf16.stride(1),
+ b_ps_u8.stride(0),
+ b_ps_u8.stride(1),
+ 0 if y_pp is None else y_pp.stride(0),
+ y.stride(0) if y_pp is None else y_pp.stride(1),
+ y.stride(1) if y_pp is None else y_pp.stride(2),
+ s_ps_u8.stride(0),
+ s_ps_u8.stride(1),
+ PREQUANT=True,
+ **cfg,
+ )
+
+ if y_pp is not None:
+ actual_ksplit = int(_TRITON.cdiv(runtime_k, int(cfg["SPLITK_BLOCK_SIZE"]) // 2))
+ grid_reduce = (_ceil_div(m, 16), _ceil_div(runtime_n, 16))
+ _REDUCE_KERNEL[grid_reduce](
+ y_pp,
+ y,
+ m,
+ runtime_n,
+ y_pp.stride(0),
+ y_pp.stride(1),
+ y_pp.stride(2),
+ y.stride(0),
+ y.stride(1),
+ 16,
+ 16,
+ actual_ksplit,
+ int(_TRITON.next_power_of_2(int(cfg["NUM_KSPLIT"]))),
+ )
+ return y
+ except Exception:
+ _DIRECT_KERNEL_SHAPE_SUPPORT[shape] = False
+ raise
+
+
+ def _run_direct_helper_path(
+ a_bf16: torch.Tensor,
+ b_shuffle: torch.Tensor,
+ b_scale_sh: torch.Tensor,
+ m: int,
+ n: int,
+ k: int,
+ ) -> torch.Tensor:
+ _resolve_runtime()
if _DIRECT_HELPER is None:
- raise RuntimeError("a16 preshuffle helper not available")
+ raise RuntimeError("direct helper unavailable")
shape = (m, n, k)
- if not _DIRECT_SHAPE_SUPPORT.get(shape, True):
+ if not _DIRECT_HELPER_SHAPE_SUPPORT.get(shape, True):
raise RuntimeError(f"direct helper disabled for {shape}")
- entry = _pick_shape_entry(m, n, k)
- b_ps, s_ps = _get_cached_preshuffle_views(b_shuffle, b_scale_sh, n, k)
- cfg = entry["cfg"]
+ cfg = _prepare_helper_cfg(m, n, k)
+ b_ps_u8, s_ps_u8 = _get_cached_preshuffle_views(b_shuffle, b_scale_sh, n, k)
+ y = _get_cached_output(a_bf16.device, m, n)
try:
config_arg = cfg
if not _DIRECT_HELPER_ACCEPTS_DICT and _SERIALIZE_DICT is not None:
- config_arg = _SERIALIZE_DICT(_config_to_dict(cfg))
+ config_arg = _SERIALIZE_DICT(cfg)
return _DIRECT_HELPER(
a_bf16,
- b_ps,
- s_ps,
+ b_ps_u8,
+ s_ps_u8,
prequant=True,
dtype=torch.bfloat16,
- y=None,
+ y=y,
config=config_arg,
skip_reduce=False,
)
except Exception:
- _DIRECT_SHAPE_SUPPORT[shape] = False
+ _DIRECT_HELPER_SHAPE_SUPPORT[shape] = False
raise
⋯ 6 unchanged lines
n: int,
k: int,
) -> torch.Tensor:
- shape = (m, n, k)
- kernel_config = ASM_KERNEL_CONFIGS.get(shape)
- if ASM_GEMM is not None and kernel_config is not None:
- kernel_name, log2_k_split = kernel_config
- if _ASM_SHAPE_SUPPORT.get(shape, True):
- out = torch.empty((m, n), dtype=torch.bfloat16, device=a_q.device)
- try:
- ASM_GEMM(
- a_q,
- b_shuffle,
- a_scale_sh,
- b_scale_sh,
- out,
- kernelName=kernel_name,
- bpreshuffle=True,
- log2_k_split=log2_k_split,
- )
- _ASM_SHAPE_SUPPORT[shape] = True
- return out
- except Exception:
- _ASM_SHAPE_SUPPORT[shape] = False
-
return aiter.gemm_a4w4(
a_q,
b_shuffle,
⋯ 12 unchanged lines
m, k = A.shape
n = B_shuffle.shape[0]
+ shape = (m, n, k)
+ helper_cfg = _prepare_helper_cfg(m, n, k)
+ b_ps_u8, _s_ps_u8 = _get_cached_preshuffle_views(B_shuffle, B_scale_sh, n, k)
+ runtime_n = b_ps_u8.shape[0] * 16
+ runtime_k = b_ps_u8.shape[1] // 16
+ direct_cfg = _prepare_direct_cfg(m, n, k, runtime_k)
try:
- return _run_direct_session_path(A, B_shuffle, B_scale_sh, m, n, k)
- except Exception:
+ out = _run_direct_kernel_path(A, B_shuffle, B_scale_sh, m, n, k)
+ _emit_path(shape, "direct", direct_cfg, f"rn={runtime_n},rk={runtime_k}")
+ return out
+ except Exception as direct_exc:
+ direct_detail = repr(direct_exc)
+
+ try:
+ out = _run_direct_helper_path(A, B_shuffle, B_scale_sh, m, n, k)
+ _emit_path(shape, "helper", helper_cfg, f"rn={runtime_n},rk={runtime_k},direct={direct_detail}")
+ return out
+ except Exception as helper_exc:
A_q, A_scale_sh = _get_cached_a_quant(A)
+ _emit_path(
+ shape,
+ "fallback",
+ helper_cfg,
+ f"rn={runtime_n},rk={runtime_k},direct={direct_detail},helper={repr(helper_exc)}",
+ )
return _run_fallback_gemm(A_q, B_shuffle, A_scale_sh, B_scale_sh, m, n, k)
scrolls · 580 diff lines total

Best evidence level for this revision: reported

JSON