Skip to content
KernelIndex
Search⌘K

submission 589522

wsxhjnb1 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-589522?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
167.2µs
#293 of 782
2026-03-19

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:2cc44370147715f9021a9511123e416b7cc07843b4fb8cbcda8d840db7a80c6c
license declaredunknown
license concludedunknown
authorswsxhjnb1
imported2026-08-15

Techniques

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

fp4Local submission entrypoint for the GPU MODE `moe-mxfp4` challenge.
split-ksplitk=1,

Kernel source

submission.py387 lines
"""
Local submission entrypoint for the GPU MODE `moe-mxfp4` challenge.

This keeps the official interface intact while making the failure mode clearer
when the ROCm/AITER stack is not present in the local Python environment.
"""

import csv
import os
import tempfile
from pathlib import Path

import torch

from task import input_t, output_t

_RUNTIME_READY = False
_OFFICIAL_E33_SHAPES = (
    (16, 512, 33, 32),
    (128, 512, 33, 64),
)


def _candidate_fmoe_configs(pkg_dir: Path) -> list[Path]:
    configs_dir = pkg_dir / "configs"
    model_dir = configs_dir / "model_configs"
    paths = []

    default_cfg = configs_dir / "tuned_fmoe.csv"
    if default_cfg.is_file():
        paths.append(default_cfg)

    if model_dir.is_dir():
        paths.extend(
            sorted(
                p
                for p in model_dir.glob("*tuned_fmoe*.csv")
                if p.is_file() and "untuned" not in p.name
            )
        )

    return paths


def _build_extra_fmoe_rows(model_dir: Path) -> list[dict[str, str]]:
    dsv3_cfg = model_dir / "dsv3_fp4_tuned_fmoe.csv"
    if not dsv3_cfg.is_file():
        return []

    with dsv3_cfg.open(newline="") as handle:
        rows = list(csv.DictReader(handle))

    extra_rows = []
    template_row = None

    # `fused_moe` keys on padded token count exactly. The bundled DeepSeek-V3 FP4
    # table starts at 32 tokens, while this challenge also benchmarks bs=16.
    for row in rows:
        if row.get("_tag", ""):
            continue
        if (
            row["token"] == "32"
            and row["model_dim"] == "7168"
            and row["inter_dim"] == "256"
            and row["expert"] == "257"
            and row["topk"] == "9"
            and row["act_type"] == "ActivationType.Silu"
            and row["dtype"] == "torch.bfloat16"
            and row["q_type"] == "QuantType.per_1x32"
        ):
            template_row = row.copy()
            cloned = row.copy()
            cloned["token"] = "16"
            extra_rows.append(cloned)
            break

    if template_row is None:
        return extra_rows

    # The official E=33 cases do not have tuned per-shape fmoe config rows in
    # AITER. Try a modest split-K override while keeping the heuristic block_m.
    for token, inter_dim, expert, block_m in _OFFICIAL_E33_SHAPES:
        row = template_row.copy()
        row.update(
            {
                "token": str(token),
                "inter_dim": str(inter_dim),
                "expert": str(expert),
                "block_m": str(block_m),
                "ksplit": "2",
                "us1": "0.0",
                "kernelName1": "",
                "err1": "0.0%",
                "us2": "0.0",
                "kernelName2": "",
                "err2": "0.0%",
                "us": "0.0",
                "run_1stage": "0",
                "tflops": "0.0",
                "bw": "0.0",
                "_tag": "",
            }
        )
        extra_rows.append(row)

    return extra_rows


def _prepare_runtime():
    global _RUNTIME_READY
    if _RUNTIME_READY:
        return

    # Several high-ranking public submissions hint that forcing NT loads helps
    # the larger expert-parallel benchmark cases.
    os.environ.setdefault("AITER_USE_NT", "1")
    os.environ.setdefault("AITER_USE_OPUS_MOE_SORTING", "0")

    try:
        import aiter
    except ImportError:
        _RUNTIME_READY = True
        return

    pkg_dir = Path(aiter.__file__).resolve().parent
    config_paths = _candidate_fmoe_configs(pkg_dir)
    extra_rows = _build_extra_fmoe_rows(pkg_dir / "configs" / "model_configs")

    if extra_rows:
        extra_cfg = Path(tempfile.gettempdir()) / "amd_moe_mxfp4_extra_fmoe.csv"
        with extra_cfg.open("w", newline="") as handle:
            writer = csv.DictWriter(handle, fieldnames=list(extra_rows[0].keys()))
            writer.writeheader()
            writer.writerows(extra_rows)
        config_paths.insert(0, extra_cfg)

    if config_paths:
        os.environ["AITER_CONFIG_FMOE"] = os.pathsep.join(str(path) for path in config_paths)

    _RUNTIME_READY = True


def _load_aiter():
    _prepare_runtime()
    try:
        from aiter import ActivationType, QuantType, dtypes
        from aiter.fused_moe import (
            fused_moe,
            get_block_size_M,
            moe_sorting,
            use_nt,
        )
        from aiter.ops.moe_op import ck_moe_stage1_fwd, ck_moe_stage2_fwd
        from aiter.ops.triton.quant.fused_mxfp4_quant import (
            fused_dynamic_mxfp4_quant_moe_sort,
        )
    except ImportError as exc:
        raise RuntimeError(
            "This moe-mxfp4 baseline uses AITER's fused_moe kernel. "
            "Install an AMD/ROCm environment with `aiter` available before "
            "running local correctness or benchmark jobs."
        ) from exc
    return (
        ActivationType,
        QuantType,
        dtypes,
        fused_moe,
        moe_sorting,
        get_block_size_M,
        use_nt,
        fused_dynamic_mxfp4_quant_moe_sort,
        ck_moe_stage1_fwd,
        ck_moe_stage2_fwd,
    )


def _should_try_direct_ck(config: dict, topk_ids) -> bool:
    return (
        config.get("bs") == 512
        and config.get("n_routed_experts") == 32
        and config.get("n_shared_experts") == 1
        and config.get("total_top_k") == 9
        and topk_ids.shape[1] == 9
    )


def _run_direct_ck_mxfp4(
    hidden_states,
    gate_up_weight_shuffled,
    down_weight_shuffled,
    gate_up_weight_scale_shuffled,
    down_weight_scale_shuffled,
    topk_weights,
    topk_ids,
    ActivationType,
    QuantType,
    dtypes,
    moe_sorting,
    get_block_size_M,
    use_nt,
    fused_dynamic_mxfp4_quant_moe_sort,
    ck_moe_stage1_fwd,
    ck_moe_stage2_fwd,
):
    token_num = hidden_states.shape[0]
    topk = topk_ids.shape[1]
    expert = gate_up_weight_shuffled.shape[0]
    model_dim = down_weight_shuffled.shape[1]
    inter_dim = down_weight_shuffled.shape[2] * 2
    block_size_M = get_block_size_M(token_num, topk, expert, inter_dim)
    use_non_temporal_load = use_nt(token_num, topk, expert)

    sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_out = moe_sorting(
        topk_ids,
        topk_weights,
        expert,
        model_dim,
        hidden_states.dtype,
        block_size_M,
        None,
        None,
        0,
    )

    a1, a1_scale = fused_dynamic_mxfp4_quant_moe_sort(
        hidden_states,
        sorted_ids=sorted_ids,
        num_valid_ids=num_valid_ids,
        token_num=token_num,
        topk=1,
        block_size=block_size_M,
    )

    a2 = torch.empty(
        (token_num, topk, inter_dim),
        dtype=hidden_states.dtype,
        device=hidden_states.device,
    )
    a2 = ck_moe_stage1_fwd(
        a1,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        sorted_ids,
        sorted_expert_ids,
        num_valid_ids,
        a2,
        topk,
        kernelName="",
        w1_scale=gate_up_weight_scale_shuffled.view(dtypes.fp8_e8m0),
        a1_scale=a1_scale,
        block_m=block_size_M,
        sorted_weights=None,
        quant_type=QuantType.per_1x32,
        activation=ActivationType.Silu,
        splitk=1,
        use_non_temporal_load=use_non_temporal_load,
        dst_type=hidden_states.dtype,
    )

    a2 = a2.view(-1, inter_dim)
    a2, a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
        a2,
        sorted_ids=sorted_ids,
        num_valid_ids=num_valid_ids,
        token_num=token_num,
        topk=topk,
        block_size=block_size_M,
    )
    a2 = a2.view(token_num, topk, -1)

    ck_moe_stage2_fwd(
        a2,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        sorted_ids,
        sorted_expert_ids,
        num_valid_ids,
        moe_out,
        topk,
        kernelName="",
        w2_scale=down_weight_scale_shuffled.view(dtypes.fp8_e8m0),
        a2_scale=a2_scale,
        block_m=block_size_M,
        sorted_weights=sorted_weights,
        quant_type=QuantType.per_1x32,
        activation=ActivationType.Silu,
        use_non_temporal_load=use_non_temporal_load,
    )

    return moe_out


def custom_kernel(data: input_t) -> output_t:
    """
    Submission template for DeepSeek-R1 MXFP4 MoE kernel.

    Input data tuple:
        hidden_states:                [M, d_hidden]                           bf16
        gate_up_weight:               [E, 2*d_expert_pad, d_hidden_pad//2]    fp4x2  (raw)
        down_weight:                  [E, d_hidden_pad, d_expert_pad//2]      fp4x2  (raw)
        gate_up_weight_scale:         [E, 2*d_expert_pad, scale_K]            e8m0   (raw)
        down_weight_scale:            [E, d_hidden_pad, scale_K]              e8m0   (raw)
        gate_up_weight_shuffled:      [E, 2*d_expert_pad, d_hidden_pad//2]    fp4x2  (shuffled)
        down_weight_shuffled:         [E, d_hidden_pad, d_expert_pad//2]      fp4x2  (shuffled)
        gate_up_weight_scale_shuffled:[padded, flat]                          e8m0   (shuffled)
        down_weight_scale_shuffled:   [padded, flat]                          e8m0   (shuffled)
        topk_weights:                 [M, total_top_k]                        float32
        topk_ids:                     [M, total_top_k]                        int32
        config:                       dict

    Returns:
        output: [M, d_hidden] bf16
    """
    (
        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

    (
        ActivationType,
        QuantType,
        dtypes,
        fused_moe,
        moe_sorting,
        get_block_size_M,
        use_nt,
        fused_dynamic_mxfp4_quant_moe_sort,
        ck_moe_stage1_fwd,
        ck_moe_stage2_fwd,
    ) = _load_aiter()

    hidden_pad = max(config["d_hidden_pad"] - config["d_hidden"], 0)
    intermediate_pad = max(config["d_expert_pad"] - config["d_expert"], 0)

    if _should_try_direct_ck(config, topk_ids):
        try:
            return _run_direct_ck_mxfp4(
                hidden_states,
                gate_up_weight_shuffled,
                down_weight_shuffled,
                gate_up_weight_scale_shuffled,
                down_weight_scale_shuffled,
                topk_weights,
                topk_ids,
                ActivationType,
                QuantType,
                dtypes,
                moe_sorting,
                get_block_size_M,
                use_nt,
                fused_dynamic_mxfp4_quant_moe_sort,
                ck_moe_stage1_fwd,
                ck_moe_stage2_fwd,
            )
        except Exception:
            pass

    output = 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,
        hidden_pad=hidden_pad,
        intermediate_pad=intermediate_pad,
    )

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