Skip to content
KernelIndex
Search⌘K

submission 564099

oofbaroomf · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_543848_257bs512k1_noopus.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-564099?include=source"
interfacepython
Compatibility
measured onAMD Instinct MI355X
declared hardwareAMD Instinct MI355X
architecturesgfx950
dtypesbf16, fp32, fp8_e8m0, int32, mxfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
AMD MXFP4 MoEsuite of 7 cases
AMD Instinct MI355X
144.4µs
#130 of 782
2026-03-15

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:b289f11bccbc6ac288f3daded1449241360ab8b75658a2dde52855203f1b1b28
license declaredunknown
license concludedunknown
authorsoofbaroomf
imported2026-08-15

Techniques

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

split-ksplit_k=ksplit,

Kernel source

submission_543848_257bs512k1_noopus.py620 lines
import functools
import importlib.util
import inspect
from pathlib import Path

import torch
from task import input_t, output_t


def _set_fmoe_tune_path() -> None:
    spec = importlib.util.find_spec("aiter")
    if spec is None or spec.origin is None:
        return
    aiter_root = Path(spec.origin).resolve().parent
    tune_file = aiter_root / "configs" / "model_configs" / "dsv3_fp4_tuned_fmoe.csv"
    if tune_file.exists():
        import os

        os.environ["AITER_CONFIG_FMOE"] = str(tune_file)


_set_fmoe_tune_path()

import aiter  # noqa: E402
import aiter.fused_moe as _fused_moe_mod  # noqa: E402
from aiter import ActivationType, QuantType  # noqa: E402


_QSORT_THRESHOLD = 256
_FORCE_SEP_QSORT_257_TOKENS = (16, 128)
_KERNEL_33_512_STAGE1 = (
    "moe_ck2stages_gemm1_256x64x128x128_1x4_"
    "MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_"
    "FP4X2_FP4X2_B16"
)
_KERNEL_33_512_STAGE2 = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce"
_KERNEL_33_2048_STAGE1 = (
    "moe_ck2stages_gemm1_256x128x128x128_1x4_"
    "MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_"
    "FP4X2_FP4X2_B16"
)
_KERNEL_33_2048_STAGE2 = (
    "moe_ck2stages_gemm2_256x128x128x128_1x4_"
    "MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_"
    "FP4X2_FP4X2_B16"
)
_A2_CACHE: dict[tuple[str, tuple[int, ...], str], torch.Tensor] = {}
_SORT_CACHE: dict[
    tuple[str, int, int, int, int, str, int],
    tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor],
] = {}


def _pad_rows(x: torch.Tensor, target_rows: int) -> torch.Tensor:
    if x.shape[0] >= target_rows:
        return x
    pad_rows = target_rows - x.shape[0]
    pad = x[-1:].expand(pad_rows, *x.shape[1:]).clone()
    return torch.cat((x, pad), dim=0)


def _repartial_with(stage, **updates):
    fn = getattr(stage, "func", None)
    if fn is None:
        return stage
    params = inspect.signature(fn).parameters
    kwargs = dict(stage.keywords or {})
    for key, value in updates.items():
        if key in params:
            kwargs[key] = value
    return functools.partial(fn, *(stage.args or ()), **kwargs)


def _make_k2_metadata(*args):
    (
        token,
        _model_dim,
        _inter_dim,
        _expert,
        _topk,
        _dtype,
        _q_dtype_a,
        _q_dtype_w,
        _q_type,
        use_g1u1,
        activation,
        _doweight_stage1,
        hidden_pad,
        intermediate_pad,
        _is_shuffled,
    ) = args
    ksplit = 2
    return _fused_moe_mod.MOEMetadata(
        functools.partial(
            _fused_moe_mod.cktile_moe_stage1,
            n_pad_zeros=intermediate_pad // 64 * 64 * (2 if use_g1u1 else 1),
            k_pad_zeros=hidden_pad // 128 * 128,
            activation=activation,
            split_k=ksplit,
        ),
        functools.partial(
            _fused_moe_mod.cktile_moe_stage2,
            n_pad_zeros=hidden_pad // 64 * 64,
            k_pad_zeros=intermediate_pad // 128 * 128,
            activation=activation,
        ),
        16 if token < 2048 else 32 if token < 16384 else 64,
        ksplit,
        False,
        False,
        True,
    )


def _make_exact_ck_metadata(*args, block_m: int, kernel1: str, kernel2: str):
    (
        _token,
        _model_dim,
        _inter_dim,
        _expert,
        _topk,
        dtype,
        _q_dtype_a,
        _q_dtype_w,
        q_type,
        _use_g1u1,
        activation,
        _doweight_stage1,
        _hidden_pad,
        _intermediate_pad,
        _is_shuffled,
    ) = args
    if kernel2.startswith("flydsl_"):
        stage2 = functools.partial(
            _fused_moe_mod._flydsl_stage2_wrapper,
            kernelName=kernel2,
        )
    else:
        stage2 = functools.partial(
            aiter.ck_moe_stage2_fwd,
            kernelName=kernel2,
            activation=activation,
            quant_type=q_type,
            use_non_temporal_load=True,
        )
    return _fused_moe_mod.MOEMetadata(
        functools.partial(
            _fused_moe_mod.ck_moe_stage1,
            kernelName=kernel1,
            activation=activation,
            quant_type=q_type,
            splitk=0,
            use_non_temporal_load=True,
            dtype=dtype,
        ),
        stage2,
        block_m,
        0,
        False,
        False,
        True,
    )


def _make_k1_large_metadata(md):
    if md.run_1stage:
        return md
    stage1 = _repartial_with(md.stage1, splitk=1, split_k=1, use_non_temporal_load=True)
    stage2 = _repartial_with(md.stage2, use_non_temporal_load=True)
    return _fused_moe_mod.MOEMetadata(
        stage1=stage1,
        stage2=stage2,
        block_m=md.block_m,
        ksplit=1,
        run_1stage=md.run_1stage,
        has_bias=md.has_bias,
        use_non_temporal_load=True,
    )


def _patch_metadata() -> None:
    if getattr(
        _fused_moe_mod,
        "_oof_force_nt_sep256_257k2_33d512flydsl_33d2048bm64",
        False,
    ):
        return

    orig_get_2stage_cfgs = _fused_moe_mod.get_2stage_cfgs

    @functools.lru_cache(maxsize=2048)
    def wrapped_get_2stage_cfgs(*args):
        (
            token,
            _model_dim,
            inter_dim,
            expert,
            _topk,
            _dtype,
            _q_dtype_a,
            q_dtype_w,
            q_type,
            _use_g1u1,
            _activation,
            _doweight_stage1,
            _hidden_pad,
            _intermediate_pad,
            is_shuffled,
        ) = args
        if (
            expert == 257
            and inter_dim == 256
            and token <= 128
            and q_type == QuantType.per_1x32
            and q_dtype_w == aiter.dtypes.fp4x2
            and is_shuffled
        ):
            return _make_k2_metadata(*args)
        if (
            expert == 33
            and inter_dim == 512
            and token <= 128
            and q_type == QuantType.per_1x32
            and q_dtype_w == aiter.dtypes.fp4x2
            and is_shuffled
        ):
            return _make_k2_metadata(*args)
        if (
            expert == 33
            and inter_dim == 512
            and token == 512
            and q_type == QuantType.per_1x32
            and q_dtype_w == aiter.dtypes.fp4x2
            and is_shuffled
        ):
            return _make_exact_ck_metadata(
                *args,
                block_m=64,
                kernel1=_KERNEL_33_512_STAGE1,
                kernel2=_KERNEL_33_512_STAGE2,
            )
        md = orig_get_2stage_cfgs(*args)
        if (
            expert == 257
            and inter_dim == 256
            and token == 512
            and q_type == QuantType.per_1x32
            and q_dtype_w == aiter.dtypes.fp4x2
            and is_shuffled
        ):
            return _make_k1_large_metadata(md)
        if (
            expert == 33
            and inter_dim == 2048
            and token == 512
            and q_type == QuantType.per_1x32
            and q_dtype_w == aiter.dtypes.fp4x2
            and is_shuffled
        ):
            return _make_exact_ck_metadata(
                *args,
                block_m=32,
                kernel1=_KERNEL_33_2048_STAGE1,
                kernel2=_KERNEL_33_2048_STAGE2,
            )
        if md.run_1stage:
            return md
        stage1 = _repartial_with(md.stage1, use_non_temporal_load=True)
        stage2 = _repartial_with(md.stage2, use_non_temporal_load=True)
        return _fused_moe_mod.MOEMetadata(
            stage1=stage1,
            stage2=stage2,
            block_m=md.block_m,
            ksplit=md.ksplit,
            run_1stage=md.run_1stage,
            has_bias=md.has_bias,
            use_non_temporal_load=True,
        )

    _fused_moe_mod.get_2stage_cfgs = wrapped_get_2stage_cfgs
    _fused_moe_mod._oof_force_nt_sep256_257k2_33d512flydsl_33d2048bm64 = True


def _patch_qsort_threshold() -> None:
    if getattr(_fused_moe_mod, "_oof_qsort_threshold", None) == _QSORT_THRESHOLD:
        return
    source = inspect.getsource(_fused_moe_mod.fused_moe_2stages)
    source = source.replace(
        "token_num_quant_moe_sort_switch = 1024",
        f"token_num_quant_moe_sort_switch = {_QSORT_THRESHOLD}",
        1,
    )
    source = source.replace(
        "    E, model_dim, inter_dim = get_inter_dim(w1.shape, w2.shape)",
        "    E, model_dim, inter_dim = get_inter_dim(w1.shape, w2.shape)\n"
        "    force_sep_qsort_257 = (\n"
        f"        token_num in {_FORCE_SEP_QSORT_257_TOKENS}\n"
        "        and E == 257\n"
        "        and inter_dim == 256\n"
        "    )",
        1,
    )
    source = source.replace(
        "if token_num <= token_num_quant_moe_sort_switch:",
        "if token_num <= token_num_quant_moe_sort_switch and not force_sep_qsort_257:",
        2,
    )
    source = source.replace("a2 = torch.empty(", "a2 = _oof_get_cached_empty(", 2)
    _fused_moe_mod.__dict__["_oof_get_cached_empty"] = _get_cached_empty
    namespace: dict[str, object] = {}
    exec(source, _fused_moe_mod.__dict__, namespace)
    _fused_moe_mod.fused_moe_2stages = namespace["fused_moe_2stages"]
    _fused_moe_mod._oof_qsort_threshold = _QSORT_THRESHOLD


def _get_cached_empty(*args, **kwargs):
    shape = args[0] if args else kwargs.get("size")
    try:
        shape = tuple(shape)
    except TypeError:
        return torch.empty(*args, **kwargs)
    if len(shape) != 3 or shape[0] != 512:
        return torch.empty(*args, **kwargs)
    dtype = kwargs.get("dtype")
    device = kwargs.get("device")
    key = (str(device), shape, str(dtype))
    cached = _A2_CACHE.get(key)
    if cached is None:
        cached = torch.empty(*args, **kwargs)
        _A2_CACHE[key] = cached
    return cached


def _get_cached_sort_buffers(
    topk_ids: torch.Tensor,
    num_experts: int,
    model_dim: int,
    moebuf_dtype: torch.dtype,
    block_size: int,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor] | None:
    m, topk = topk_ids.shape
    if m != 512:
        return None
    device = topk_ids.device
    max_num_tokens_padded = int(topk_ids.numel() + num_experts * block_size - topk)
    max_num_m_blocks = int((max_num_tokens_padded + block_size - 1) // block_size)
    key = (
        str(device),
        m,
        topk,
        int(num_experts),
        int(model_dim),
        str(moebuf_dtype),
        int(block_size),
    )
    cached = _SORT_CACHE.get(key)
    if cached is None:
        cached = (
            torch.empty(max_num_tokens_padded, dtype=aiter.dtypes.i32, device=device),
            torch.empty(max_num_tokens_padded, dtype=aiter.dtypes.fp32, device=device),
            torch.empty(max_num_m_blocks, dtype=aiter.dtypes.i32, device=device),
            torch.empty(2, dtype=aiter.dtypes.i32, device=device),
            torch.empty((m, model_dim), dtype=moebuf_dtype, device=device),
        )
        _SORT_CACHE[key] = cached
    return cached


def _patch_sort_cache() -> None:
    if getattr(_fused_moe_mod, "_oof_bs512_sort_cache", False):
        return

    orig_moe_sorting_impl = _fused_moe_mod._moe_sorting_impl

    def wrapped_moe_sorting_impl(
        topk_ids,
        topk_weights,
        num_experts,
        model_dim,
        moebuf_dtype,
        block_size,
        expert_mask,
        num_local_tokens,
        dispatch_policy,
        use_opus,
    ):
        cached = _get_cached_sort_buffers(
            topk_ids, num_experts, model_dim, moebuf_dtype, block_size
        )
        if cached is None:
            return orig_moe_sorting_impl(
                topk_ids,
                topk_weights,
                num_experts,
                model_dim,
                moebuf_dtype,
                block_size,
                expert_mask,
                num_local_tokens,
                dispatch_policy,
                use_opus,
            )
        sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = cached
        fwd_fn = aiter.moe_sorting_opus_fwd if use_opus else aiter.moe_sorting_fwd
        fwd_fn(
            topk_ids,
            topk_weights,
            sorted_ids,
            sorted_weights,
            sorted_expert_ids,
            num_valid_ids,
            moe_buf,
            num_experts,
            int(block_size),
            expert_mask,
            num_local_tokens,
            dispatch_policy,
        )
        return sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf

    _fused_moe_mod._moe_sorting_impl = wrapped_moe_sorting_impl
    _fused_moe_mod._oof_bs512_sort_cache = True


def _patch_257_bs512_opus_sorting() -> None:
    if getattr(_fused_moe_mod, "_oof_sorting_patch_257_bs512_opus", False):
        return

    def wrapped_moe_sorting(
        topk_ids,
        topk_weights,
        num_experts,
        model_dim,
        moebuf_dtype,
        block_size=_fused_moe_mod.BLOCK_SIZE_M,
        expert_mask=None,
        num_local_tokens=None,
        dispatch_policy=0,
    ):
        use_opus = num_experts == 257 and topk_ids.shape[0] == 512
        return _fused_moe_mod._moe_sorting_impl(
            topk_ids,
            topk_weights,
            num_experts,
            model_dim,
            moebuf_dtype,
            block_size,
            expert_mask,
            num_local_tokens,
            dispatch_policy,
            use_opus=use_opus,
        )

    _fused_moe_mod.moe_sorting = wrapped_moe_sorting
    _fused_moe_mod._oof_sorting_patch_257_bs512_opus = True


def _patch_33_bs512_opus_sorting() -> None:
    if getattr(_fused_moe_mod, "_oof_sorting_patch_33_bs512_opus", False):
        return

    def wrapped_moe_sorting(
        topk_ids,
        topk_weights,
        num_experts,
        model_dim,
        moebuf_dtype,
        block_size=_fused_moe_mod.BLOCK_SIZE_M,
        expert_mask=None,
        num_local_tokens=None,
        dispatch_policy=0,
    ):
        use_opus = num_experts in (33, 257) and topk_ids.shape[0] == 512
        return _fused_moe_mod._moe_sorting_impl(
            topk_ids,
            topk_weights,
            num_experts,
            model_dim,
            moebuf_dtype,
            block_size,
            expert_mask,
            num_local_tokens,
            dispatch_policy,
            use_opus=use_opus,
        )

    _fused_moe_mod.moe_sorting = wrapped_moe_sorting
    _fused_moe_mod._oof_sorting_patch_33_bs512_opus = True


def _patch_cktile_stage1_out() -> None:
    if getattr(_fused_moe_mod, "_oof_cktile_stage1_use_out", False):
        return

    def wrapped_cktile_moe_stage1(
        hidden_states,
        w1,
        w2,
        sorted_token_ids,
        sorted_expert_ids,
        num_valid_ids,
        out,
        topk,
        block_m,
        a1_scale,
        w1_scale,
        sorted_weights=None,
        n_pad_zeros=0,
        k_pad_zeros=0,
        bias1=None,
        activation=ActivationType.Silu,
        split_k=1,
        dtype=torch.bfloat16,
    ):
        token_num = hidden_states.shape[0]
        _, _, k1 = w1.shape
        _, k2, n2 = w2.shape
        d = n2 if k2 == k1 else n2 * 2
        if w1.dtype is torch.uint32:
            d *= 8
        if out.shape != (token_num, topk, d) or out.dtype != dtype or out.device != hidden_states.device:
            out = torch.empty((token_num, topk, d), dtype=dtype, device=hidden_states.device)
        tmp_out = (
            torch.zeros(
                (token_num, topk, w1.shape[1]),
                dtype=hidden_states.dtype,
                device=out.device,
            )
            if split_k > 1
            else out
        )
        aiter.moe_cktile2stages_gemm1(
            hidden_states,
            w1,
            tmp_out,
            sorted_token_ids,
            sorted_expert_ids,
            num_valid_ids,
            topk,
            n_pad_zeros,
            k_pad_zeros,
            sorted_weights,
            a1_scale,
            w1_scale,
            bias1,
            activation,
            block_m,
            split_k,
        )
        if split_k > 1:
            if activation == ActivationType.Silu:
                aiter.silu_and_mul(out, tmp_out)
            else:
                aiter.gelu_and_mul(out, tmp_out)
        return out

    _fused_moe_mod.cktile_moe_stage1 = wrapped_cktile_moe_stage1
    _fused_moe_mod._oof_cktile_stage1_use_out = True


_patch_metadata()
_patch_qsort_threshold()
_patch_sort_cache()
_patch_257_bs512_opus_sorting()
_patch_33_bs512_opus_sorting()
_patch_cktile_stage1_out()


def _pick_block_size(config: dict) -> int | None:
    total_experts = config["n_routed_experts"] + config["n_shared_experts"]
    if total_experts == 33 and config["d_expert"] == 2048 and config["bs"] >= 512:
        return 64
    return None


def custom_kernel(data: input_t) -> output_t:
    (
        hidden_states,
        _gate_up_weight,
        _down_weight,
        _gate_up_weight_scale,
        _down_weight_scale,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        gate_up_weight_scale_shuffled,
        down_weight_scale_shuffled,
        topk_weights,
        topk_ids,
        config,
    ) = data

    hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
    intermediate_pad = config["d_expert_pad"] - config["d_expert"]
    use_shape7_exact_ck = (
        config["n_routed_experts"] + config["n_shared_experts"] == 33
        and config["d_expert"] == 2048
        and config["bs"] == 512
    )

    output = _fused_moe_mod.fused_moe(
        hidden_states,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        topk_weights,
        topk_ids,
        expert_mask=None,
        activation=ActivationType.Silu,
        quant_type=QuantType.per_1x32,
        doweight_stage1=False,
        w1_scale=gate_up_weight_scale_shuffled,
        w2_scale=down_weight_scale_shuffled,
        a1_scale=None,
        a2_scale=None,
        block_size_M=None if use_shape7_exact_ck else _pick_block_size(config),
        hidden_pad=hidden_pad,
        intermediate_pad=intermediate_pad,
    )

    return output
scrolls · 620 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