Skip to content
KernelIndex
Search⌘K

submission 678941

.jonnss · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

Submission_v164.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-mxfp4-mm-678941?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
10.6µs
#279 of 1143
2026-03-31

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:cf0e339de7fd59f261d71a9e274fe243f0a65f734a8d331fff7d2426b0e0bf0a
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"SPLITK_BLOCK_SIZE": max(k // max(ksplit, 1), 64),

Kernel source

Submission_v164.py389 lines
"""
SESSION.md-guided A16 preshuffle rebuild.

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.
"""
import importlib
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 = 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,
    ],
] = {}

_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]]] = {}

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): ("_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_256x128E", None),
}
_FORCE_ASM_SHAPES = {
    (256, 3072, 1536),
}

_ASM_SHAPE_SUPPORT: dict[tuple[int, int, int], bool] = {}


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 _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:
        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:
            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)))

    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 * 1_000_000 + 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
        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, s_ps

    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()

    _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,
    )

    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 _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_direct_helper() -> None:
    global _DIRECT_INIT_DONE, _DIRECT_HELPER, _DIRECT_HELPER_ACCEPTS_DICT, _SERIALIZE_DICT
    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

    candidates: list[tuple[object, str, bool]] = [
        (aiter, "gemm_a16wfp4_preshuffle", True),
        (aiter, "gemm_a16wfp4_preshuffle_", False),
    ]
    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", True),
                (mod, "gemm_a16wfp4_preshuffle_", False),
            ]
        )

    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


def _pick_shape_entry(m: int, n: int, k: int) -> dict[str, 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
    waves_per_eu = 2 if wgs > _CU else 1

    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": waves_per_eu,
        "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,
    }
    _SHAPE_CACHE[shape] = entry
    return entry


def _run_direct_session_path(
    a_bf16: torch.Tensor,
    b_shuffle: torch.Tensor,
    b_scale_sh: torch.Tensor,
    m: int,
    n: int,
    k: int,
) -> torch.Tensor:
    _resolve_direct_helper()
    if _DIRECT_HELPER is None:
        raise RuntimeError("a16 preshuffle helper not available")

    shape = (m, n, k)
    if not _DIRECT_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"]

    try:
        config_arg = cfg
        if not _DIRECT_HELPER_ACCEPTS_DICT and _SERIALIZE_DICT is not None:
            config_arg = _SERIALIZE_DICT(_config_to_dict(cfg))
        return _DIRECT_HELPER(
            a_bf16,
            b_ps,
            s_ps,
            prequant=True,
            dtype=torch.bfloat16,
            y=None,
            config=config_arg,
            skip_reduce=False,
        )
    except Exception:
        _DIRECT_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:
    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,
        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)

    if shape in _FORCE_ASM_SHAPES:
        A_q, A_scale_sh = _get_cached_a_quant(A)
        return _run_fallback_gemm(A_q, B_shuffle, A_scale_sh, B_scale_sh, m, n, k)

    try:
        return _run_direct_session_path(A, B_shuffle, B_scale_sh, m, n, k)
    except Exception:
        A_q, A_scale_sh = _get_cached_a_quant(A)
        return _run_fallback_gemm(A_q, B_shuffle, A_scale_sh, B_scale_sh, m, n, k)
scrolls · 389 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 529737.

- #!POPCORN leaderboard amd-mxfp4-mm
"""
- FP4 quant + FP4 GEMM reference: bf16 A, MXFP4 B -> MXFP4 per-1x32 quant A -> gemm_a4w4 -> bf16 C.
- Quant logic follows aiter op_tests/test_gemm_a4w4.py (get_triton_quant(QuantType.per_1x32)).
+ SESSION.md-guided A16 preshuffle rebuild.
+
+ 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.
"""
- import os
+ import importlib
+ 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
- VARIANT = os.environ.get("MXFP4_MM_VARIANT", "baseline")
- SPLIT_ROW_M = 32
- SPLIT_ROW_K = 512
- SPLIT_CHUNK_M = 16
+ _CU = 256
+ _LOW_UTIL_THRESHOLD = (_CU * 3) // 4
- def _run_quant_gemm(aiter, quant_func, dtypes, A: torch.Tensor, B_shuffle: torch.Tensor, B_scale_sh: torch.Tensor) -> torch.Tensor:
- A_q, A_scale_sh = quant_func(A, shuffle=True)
- return aiter.gemm_a4w4(
- A_q,
- B_shuffle,
- A_scale_sh,
- B_scale_sh,
- dtype=dtypes.bf16,
- bpreshuffle=True,
+ _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,
+ ],
+ ] = {}
+
+ _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]]] = {}
+
+ 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): ("_ZN5aiter42f4gemm_bf16_per1x32Fp4_BpreShuffle_256x128E", None),
+ }
+ _FORCE_ASM_SHAPES = {
+ (256, 3072, 1536),
+ }
+
+ _ASM_SHAPE_SUPPORT: dict[tuple[int, int, int], bool] = {}
+
+
+ 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 _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:
+ 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:
+ 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)))
+
+ 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 * 1_000_000 + 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
+ 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, s_ps
+
+ 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()
+
+ _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,
)
+ 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)))
- def _run_split_row_two_pass(
- aiter,
- quant_func,
- dtypes,
- A: torch.Tensor,
- B_shuffle: torch.Tensor,
- B_scale_sh: torch.Tensor,
+ return b_ps_u8, s_ps_u8
+
+
+ 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_direct_helper() -> None:
+ global _DIRECT_INIT_DONE, _DIRECT_HELPER, _DIRECT_HELPER_ACCEPTS_DICT, _SERIALIZE_DICT
+ 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
+
+ candidates: list[tuple[object, str, bool]] = [
+ (aiter, "gemm_a16wfp4_preshuffle", True),
+ (aiter, "gemm_a16wfp4_preshuffle_", False),
+ ]
+ 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", True),
+ (mod, "gemm_a16wfp4_preshuffle_", False),
+ ]
+ )
+
+ 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
+
+
+ def _pick_shape_entry(m: int, n: int, k: int) -> dict[str, 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
+ waves_per_eu = 2 if wgs > _CU else 1
+
+ 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": waves_per_eu,
+ "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,
+ }
+ _SHAPE_CACHE[shape] = entry
+ return entry
+
+
+ def _run_direct_session_path(
+ a_bf16: torch.Tensor,
+ b_shuffle: torch.Tensor,
+ b_scale_sh: torch.Tensor,
+ m: int,
+ n: int,
+ k: int,
) -> torch.Tensor:
- top = _run_quant_gemm(aiter, quant_func, dtypes, A[:SPLIT_CHUNK_M], B_shuffle, B_scale_sh)
- bottom = _run_quant_gemm(aiter, quant_func, dtypes, A[SPLIT_CHUNK_M:], B_shuffle, B_scale_sh)
- return torch.cat((top, bottom), dim=0)
+ _resolve_direct_helper()
+ if _DIRECT_HELPER is None:
+ raise RuntimeError("a16 preshuffle helper not available")
+ shape = (m, n, k)
+ if not _DIRECT_SHAPE_SUPPORT.get(shape, True):
+ raise RuntimeError(f"direct helper disabled for {shape}")
- def custom_kernel(data: input_t) -> output_t:
- """
- Reference: MXFP4 per-1x32 quant on A; B_shuffle, B_scale_sh from generate_input.
- gemm_a4w4 with bpreshuffle=True.
- """
- import aiter
- from aiter import QuantType, dtypes
+ 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"]
+ try:
+ config_arg = cfg
+ if not _DIRECT_HELPER_ACCEPTS_DICT and _SERIALIZE_DICT is not None:
+ config_arg = _SERIALIZE_DICT(_config_to_dict(cfg))
+ return _DIRECT_HELPER(
+ a_bf16,
+ b_ps,
+ s_ps,
+ prequant=True,
+ dtype=torch.bfloat16,
+ y=None,
+ config=config_arg,
+ skip_reduce=False,
+ )
+ except Exception:
+ _DIRECT_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:
+ 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,
+ 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
- A = A.contiguous()
+ if not A.is_contiguous():
+ A = A.contiguous()
+
m, k = A.shape
+ n = B_shuffle.shape[0]
+ shape = (m, n, k)
- quant_func = aiter.get_triton_quant(QuantType.per_1x32)
- if VARIANT == "shape_aware":
- if m == SPLIT_ROW_M and k == SPLIT_ROW_K:
- return _run_split_row_two_pass(aiter, quant_func, dtypes, A, B_shuffle, B_scale_sh)
+ if shape in _FORCE_ASM_SHAPES:
+ A_q, A_scale_sh = _get_cached_a_quant(A)
+ return _run_fallback_gemm(A_q, B_shuffle, A_scale_sh, B_scale_sh, m, n, k)
- return _run_quant_gemm(aiter, quant_func, dtypes, A, B_shuffle, B_scale_sh)
+ try:
+ return _run_direct_session_path(A, B_shuffle, B_scale_sh, m, n, k)
+ except Exception:
+ A_q, A_scale_sh = _get_cached_a_quant(A)
+ return _run_fallback_gemm(A_q, B_shuffle, A_scale_sh, B_scale_sh, m, n, k)
scrolls · 428 diff lines total

Best evidence level for this revision: reported

JSON