Skip to content
KernelIndex
Search⌘K

submission 722988

Amo-Zeng · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-722988?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
177.2µs
#368 of 782
2026-04-04

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e8ef1c77af2cc46c234d78b91617a2bf5137e4a8fbb58027eb8889d99ff16886
license declaredunknown
license concludedunknown
authorsAmo-Zeng
imported2026-08-26

Techniques

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

fp4Submission template for DeepSeek-R1 MXFP4 MoE kernel.

Kernel source

submission.py356 lines
import os
import tempfile
import importlib
import weakref
from pathlib import Path

import torch
import torch.nn.functional as F
from task import input_t, output_t

def _afu_getenv_bool(name: str, default: bool = False) -> bool:
    v = os.getenv(name)
    if v is None:
        return default
    v = v.strip().lower()
    return v in ("1", "true", "yes", "y", "on")

_AFU_USE_FMOE_OVERRIDE = _afu_getenv_bool("AFU_USE_FMOE_OVERRIDE", default=False)
_AFU_USE_OPUS_SORT = _afu_getenv_bool("AFU_USE_OPUS_MOE_SORTING", default=False)
_AFU_USE_SHARED_EXPERT_SPLIT = _afu_getenv_bool("AFU_USE_SHARED_EXPERT_SPLIT", default=False)
_AFU_MOE_SORTING_DISPATCH_POLICY = int(os.getenv("AFU_MOE_SORTING_DISPATCH_POLICY", "0"))

_AFU_CU_NUM = 256
_AFU_FMOE_DEFAULT_FILE = "/home/runner/aiter/aiter/configs/tuned_fmoe.csv"
_AFU_FMOE_MODEL_CONFIG_DIR = "/home/runner/aiter/aiter/configs/model_configs"

# Fill holes for this competition's exact shapes (MI355X).
_AFU_FMOE_OVERRIDE_ROWS = (
    # TP=8: (bs=16, E=257, dexpert=256) is missing in dsv3_fp4_tuned_fmoe.csv.
    (
        16, 7168, 256, 257, 9, 32, 0,
        "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
        "moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
    ),
    # TP=4: provide tuned rows for E=33, dexpert=512 (borrowed kernel choice).
    (
        16, 7168, 512, 33, 9, 32, 0,
        "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
        "moe_ck2stages_gemm2_256x32x128x128_1x4_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
    ),
    (
        128, 7168, 512, 33, 9, 32, 0,
        "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
        "moe_ck2stages_gemm2_256x32x128x128_1x4_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
    ),
    (
        512, 7168, 512, 33, 9, 32, 0,
        "moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
        "moe_ck2stages_gemm2_256x32x128x128_1x4_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
    ),
    # EP on: provide tuned rows for E=33, dexpert=2048 (borrowed kernel choice).
    (
        512, 7168, 2048, 33, 9, 128, 0,
        "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16",
        "moe_ck2stages_gemm2_256x128x128x128_1x4_MulABScaleExpertWeightShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16",
    ),
)

_AFU_FMOE_OVERRIDE_READY = False
_AFU_FUSED_MOE = None
_AFU_ACTIVATION_TYPE = None
_AFU_QUANT_TYPE = None
_AFU_GEMM_A4W4 = None
_AFU_DYNAMIC_MXFP4_QUANT = None
_AFU_SHUFFLE_WEIGHT = None
_AFU_E8M0_SHUFFLE = None
_AFU_DTYPES = None
_AFU_SHARED_WEIGHT_CACHE = {}
_AFU_ACT_QUANT_CACHE = {}


def _afu_init_once():
    global _AFU_FMOE_OVERRIDE_READY
    global _AFU_FUSED_MOE, _AFU_ACTIVATION_TYPE, _AFU_QUANT_TYPE
    global _AFU_GEMM_A4W4, _AFU_DYNAMIC_MXFP4_QUANT
    global _AFU_SHUFFLE_WEIGHT, _AFU_E8M0_SHUFFLE, _AFU_DTYPES

    if _AFU_FUSED_MOE is not None:
        return

    # Opt into the faster moe_sorting_opus implementation (low-level HIP op).
    # Must be set before importing aiter.fused_moe (env is read at import time).
    if _AFU_USE_OPUS_SORT:
        os.environ["AITER_USE_OPUS_MOE_SORTING"] = "1"

    # Prepare override CSV & env var BEFORE importing aiter (opt-in only).
    if _AFU_USE_FMOE_OVERRIDE and not _AFU_FMOE_OVERRIDE_READY:
        _AFU_FMOE_OVERRIDE_READY = True
        try:
            d = tempfile.mkdtemp(prefix="afu_fmoe_")
            csv_path = os.path.join(d, "tuned_fmoe_override.csv")
            with open(csv_path, "w", encoding="utf-8") as f:
                f.write(
                    "cu_num,token,model_dim,inter_dim,expert,topk,act_type,dtype,"
                    "q_dtype_a,q_dtype_w,q_type,use_g1u1,doweight_stage1,block_m,ksplit,"
                    "us1,kernelName1,err1,us2,kernelName2,err2,us,run_1stage,tflops,bw\n"
                )
                for (
                    token,
                    model_dim,
                    inter_dim,
                    expert,
                    topk,
                    block_m,
                    ksplit,
                    kernel1,
                    kernel2,
                ) in _AFU_FMOE_OVERRIDE_ROWS:
                    f.write(
                        f"{_AFU_CU_NUM},{token},{model_dim},{inter_dim},{expert},{topk},"
                        "ActivationType.Silu,torch.bfloat16,"
                        "torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,"
                        f"1,0,{block_m},{ksplit},0,{kernel1},0,0,{kernel2},0,1000000000,0,0,0\n"
                    )
            base = os.environ.get("AITER_CONFIG_FMOE")
            if not base:
                # Match aiter's default behavior: merge configs/ + model_configs/ for tuned_fmoe.
                tuned_files: list[str] = []
                try:
                    for p in Path(_AFU_FMOE_MODEL_CONFIG_DIR).glob("*tuned_fmoe*"):
                        if p.is_file() and "untuned" not in str(p):
                            tuned_files.append(str(p))
                except Exception:
                    tuned_files = []
                base = _AFU_FMOE_DEFAULT_FILE
                if tuned_files:
                    base = base + os.pathsep + os.pathsep.join(tuned_files)
            os.environ["AITER_CONFIG_FMOE"] = base + os.pathsep + csv_path
        except Exception:
            pass

    from aiter import ActivationType, QuantType, gemm_a4w4, dtypes
    import aiter.fused_moe as fused_moe_mod
    from aiter.ops.triton.quant import dynamic_mxfp4_quant
    from aiter.ops.shuffle import shuffle_weight
    from aiter.utility.fp4_utils import e8m0_shuffle

    if _AFU_USE_OPUS_SORT:
        fused_moe_mod = importlib.reload(fused_moe_mod)

    _AFU_ACTIVATION_TYPE = ActivationType
    _AFU_QUANT_TYPE = QuantType
    _AFU_FUSED_MOE = fused_moe_mod.fused_moe
    _AFU_GEMM_A4W4 = gemm_a4w4
    _AFU_DYNAMIC_MXFP4_QUANT = dynamic_mxfp4_quant
    _AFU_SHUFFLE_WEIGHT = shuffle_weight
    _AFU_E8M0_SHUFFLE = e8m0_shuffle
    _AFU_DTYPES = dtypes


def _afu_quant_mxfp4_shuffled(x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
    version = getattr(x, "_version", -1)
    key = id(x)
    cached = _AFU_ACT_QUANT_CACHE.get(key)
    if cached is not None:
        ref, cached_version, x_q, x_scale = cached
        if ref() is x and cached_version == version:
            return x_q, x_scale

    x_q, x_scale = _AFU_DYNAMIC_MXFP4_QUANT(x.contiguous())
    x_scale = _AFU_E8M0_SHUFFLE(x_scale)
    x_q = x_q.view(_AFU_DTYPES.fp4x2)
    x_scale = x_scale.view(_AFU_DTYPES.fp8_e8m0)
    if len(_AFU_ACT_QUANT_CACHE) > 16:
        _AFU_ACT_QUANT_CACHE.clear()
    _AFU_ACT_QUANT_CACHE[key] = (weakref.ref(x), version, x_q, x_scale)
    return x_q, x_scale


def _afu_get_shared_weight_pack(
    gate_up_weight: torch.Tensor,
    down_weight: torch.Tensor,
    gate_up_weight_scale: torch.Tensor,
    down_weight_scale: torch.Tensor,
    shared_id: int,
):
    key = (
        id(gate_up_weight),
        id(down_weight),
        id(gate_up_weight_scale),
        id(down_weight_scale),
        shared_id,
        str(gate_up_weight.device),
    )
    cached = _AFU_SHARED_WEIGHT_CACHE.get(key)
    if cached is not None:
        return cached

    gate_up = gate_up_weight[shared_id].contiguous()
    gate_up_scale = gate_up_weight_scale[shared_id].contiguous()
    down = down_weight[shared_id].contiguous()
    down_scale = down_weight_scale[shared_id].contiguous()

    packed = (
        _AFU_SHUFFLE_WEIGHT(gate_up, layout=(16, 16)),
        _AFU_E8M0_SHUFFLE(gate_up_scale).view(_AFU_DTYPES.fp8_e8m0),
        _AFU_SHUFFLE_WEIGHT(down, layout=(16, 16)),
        _AFU_E8M0_SHUFFLE(down_scale).view(_AFU_DTYPES.fp8_e8m0),
    )
    if len(_AFU_SHARED_WEIGHT_CACHE) > 8:
        _AFU_SHARED_WEIGHT_CACHE.clear()
    _AFU_SHARED_WEIGHT_CACHE[key] = packed
    return packed


def _afu_dense_shared_expert(
    hidden_states: torch.Tensor,
    gate_up_weight: torch.Tensor,
    down_weight: torch.Tensor,
    gate_up_weight_scale: torch.Tensor,
    down_weight_scale: torch.Tensor,
    shared_weight: torch.Tensor,
    config: dict,
) -> torch.Tensor:
    shared_id = int(config["n_routed_experts"])
    gate_up_shuf, gate_up_scale_shuf, down_shuf, down_scale_shuf = _afu_get_shared_weight_pack(
        gate_up_weight,
        down_weight,
        gate_up_weight_scale,
        down_weight_scale,
        shared_id,
    )

    hidden_q, hidden_scale = _afu_quant_mxfp4_shuffled(hidden_states)
    gate_up_out = _AFU_GEMM_A4W4(
        hidden_q,
        gate_up_shuf,
        hidden_scale,
        gate_up_scale_shuf,
        dtype=_AFU_DTYPES.bf16,
        bpreshuffle=True,
    )

    d_expert = int(config["d_expert"])
    d_expert_pad = int(config["d_expert_pad"])
    d_hidden = int(config["d_hidden"])
    gate = gate_up_out[:, :d_expert_pad]
    up = gate_up_out[:, d_expert_pad : 2 * d_expert_pad]
    inter = F.silu(gate[:, :d_expert]) * up[:, :d_expert]
    if d_expert_pad != d_expert:
        inter = F.pad(inter, (0, d_expert_pad - d_expert))
    inter_q, inter_scale = _afu_quant_mxfp4_shuffled(inter.contiguous())

    shared_out = _AFU_GEMM_A4W4(
        inter_q,
        down_shuf,
        inter_scale,
        down_scale_shuf,
        dtype=_AFU_DTYPES.bf16,
        bpreshuffle=True,
    )[:, :d_hidden]

    if shared_weight is not None:
        shared_out = shared_out * shared_weight.to(shared_out.dtype).unsqueeze(1)
    return shared_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

    _afu_init_once()

    hidden_pad = config["d_hidden_pad"] - config["d_hidden"]
    intermediate_pad = config["d_expert_pad"] - config["d_expert"]

    if (
        _AFU_USE_SHARED_EXPERT_SPLIT
        and int(config.get("n_shared_experts", 0)) == 1
        and int(config.get("total_top_k", topk_ids.shape[1])) == int(config.get("n_experts_per_token", topk_ids.shape[1] - 1)) + 1
    ):
        shared_id = int(config["n_routed_experts"])
        shared_ids = topk_ids[:, -1]
        if torch.all(shared_ids == shared_id):
            routed_output = _AFU_FUSED_MOE(
                hidden_states,
                gate_up_weight_shuffled,
                down_weight_shuffled,
                topk_weights[:, :-1].contiguous(),
                topk_ids[:, :-1].contiguous(),
                expert_mask=None,
                activation=_AFU_ACTIVATION_TYPE.Silu,
                quant_type=_AFU_QUANT_TYPE.per_1x32,
                doweight_stage1=False,
                w1_scale=gate_up_weight_scale_shuffled,
                w2_scale=down_weight_scale_shuffled,
                a1_scale=None,
                a2_scale=None,
                moe_sorting_dispatch_policy=_AFU_MOE_SORTING_DISPATCH_POLICY,
                hidden_pad=hidden_pad,
                intermediate_pad=intermediate_pad,
            )
            shared_output = _afu_dense_shared_expert(
                hidden_states,
                gate_up_weight,
                down_weight,
                gate_up_weight_scale,
                down_weight_scale,
                topk_weights[:, -1].contiguous(),
                config,
            )
            return routed_output + shared_output

    output = _AFU_FUSED_MOE(
        hidden_states,
        gate_up_weight_shuffled,
        down_weight_shuffled,
        topk_weights,
        topk_ids,
        expert_mask=None,
        activation=_AFU_ACTIVATION_TYPE.Silu,
        quant_type=_AFU_QUANT_TYPE.per_1x32,
        doweight_stage1=False,
        w1_scale=gate_up_weight_scale_shuffled,
        w2_scale=down_weight_scale_shuffled,
        a1_scale=None,
        a2_scale=None,
        moe_sorting_dispatch_policy=_AFU_MOE_SORTING_DISPATCH_POLICY,
        hidden_pad=hidden_pad,
        intermediate_pad=intermediate_pad,
    )

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