Skip to content
KernelIndex
Search⌘K

submission 647148

HorizonLiang · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_v126.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-647148?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
136.4µs
#104 of 782
2026-03-27

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:cd6eaec5eb2b6cd676fd4f83e500fddd07af7eafea0d94a1f0c743dad6ce9b44
license declaredunknown
license concludedunknown
authorsHorizonLiang
imported2026-08-15

Techniques

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

fp4if b_dtype == "fp4":

Kernel source

submission_v126.py381 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X

import copy
import functools
import importlib
import os
import re
import sys
import tempfile
import types
from pathlib import Path

from task import input_t, output_t


os.environ["AITER_USE_OPUS_MOE_SORTING"] = "1"

_CFG_READY = False
_PATCH_READY = False
_RUNTIME_FLYDSL_READY = False

_FLYDSL_STAGE1_BY_TOKEN = {
    16: "flydsl_moe1_afp4_wfp4_bf16_t16x256x256",
    64: "flydsl_moe1_afp4_wfp4_bf16_t64x256x256",
    128: "flydsl_moe1_afp4_wfp4_bf16_t128x256x256",
    256: "flydsl_moe1_afp4_wfp4_bf16_t64x256x256",
    512: "flydsl_moe1_afp4_wfp4_bf16_t128x256x256",
}

_CUSTOM_FMOE_ROWS = """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,_tag
256,16,7168,256,257,9,ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,32,2,40.0,moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16,0.0%,24.0,moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16,0.0%,64.0,0,0.0,0.0,
256,16,7168,512,33,9,ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,32,2,20.0,moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16,0.0%,20.0,moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16,0.0%,40.0,0,0.0,0.0,
256,128,7168,512,33,9,ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,32,2,50.0,moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16,0.0%,50.0,moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16,0.0%,100.0,0,0.0,0.0,
256,512,7168,512,33,9,ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,64,0,90.0,moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16,0.0%,70.0,flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce,0.1%,160.0,0,0.0,0.0,
256,512,7168,512,33,9,ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,32,0,60.0,moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16,0.0%,62.0,flydsl_moe2_afp4_wfp4_bf16_t32x128x256_reduce,0.1%,122.0,0,0.0,0.0,flydsl_fallback
256,512,7168,2048,33,9,ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,64,0,140.0,moe_ck2stages_gemm1_256x64x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16,0.0%,110.0,flydsl_moe2_afp4_wfp4_bf16_t64x128x256_reduce,0.4%,250.0,0,0.0,0.0,
"""


def _ensure_custom_fmoe_config() -> None:
    global _CFG_READY
    if _CFG_READY:
        return

    from aiter.jit.core import AITER_CONFIGS, AITER_ROOT_DIR

    cfg_dir = Path(AITER_ROOT_DIR) / "aiter" / "configs"
    custom_cfg = Path(tempfile.gettempdir()) / "aiter_configs" / "submission_v88.csv"
    custom_cfg.parent.mkdir(parents=True, exist_ok=True)
    custom_cfg.write_text(_CUSTOM_FMOE_ROWS, encoding="utf-8")

    config_paths = [
        custom_cfg,
        cfg_dir / "tuned_fmoe.csv",
        cfg_dir / "model_configs" / "a8w8_blockscale_tuned_fmoe_qwen3_235b.csv",
        cfg_dir / "model_configs" / "dsv3_fp4_tuned_fmoe.csv",
    ]
    os.environ["AITER_CONFIG_FMOE"] = os.pathsep.join(
        str(path) for path in config_paths if path.exists()
    )

    AITER_CONFIGS.get_config_file.cache_clear()
    fused_moe_module = importlib.import_module("aiter.fused_moe")
    fused_moe_module.cfg_2stages = None
    fused_moe_module.get_2stage_cfgs.cache_clear()
    _CFG_READY = True


def _ensure_flydsl_runtime_patch() -> None:
    global _RUNTIME_FLYDSL_READY
    if _RUNTIME_FLYDSL_READY:
        return

    mixed_module = importlib.import_module(
        "aiter.ops.flydsl.kernels.mixed_moe_gemm_2stage"
    )
    kernels_pkg = importlib.import_module("aiter.ops.flydsl.kernels")
    moe_kernels_module = importlib.import_module("aiter.ops.flydsl.moe_kernels")

    source_path = Path(mixed_module.__file__)
    source = source_path.read_text(encoding="utf-8")
    patched_source = source
    signature_hits = 0
    scale_hits = 0

    if "b_scale_gate=None" in source and "b_scale_up=None" in source:
        old_signature = (
            "                        b_tile_in_gate,\n"
            "                        b_tile_in_up,\n"
            "                        lds_base,"
        )
        new_signature = (
            "                        b_tile_in,\n"
            "                        lds_base,"
        )
        if old_signature in patched_source:
            patched_source = patched_source.replace(old_signature, new_signature, 1)
            signature_hits = 1

        old_scale = (
            "                        b_scale_gate=None,\n"
            "                        b_scale_up=None,"
        )
        new_scale = "                        b_scale=None,"
        if old_scale in patched_source:
            patched_source = patched_source.replace(old_scale, new_scale, 1)
            scale_hits = 1

    if patched_source == source:
        _RUNTIME_FLYDSL_READY = True
        return

    patched_dir = Path(tempfile.gettempdir()) / "submission_v88_flydsl_patch"
    patched_dir.mkdir(parents=True, exist_ok=True)
    patched_source_path = patched_dir / source_path.name
    patched_source_path.write_text(patched_source, encoding="utf-8")

    sys.stderr.write(
        "[submission_v88] wrote patched FlyDSL source to "
        f"{patched_source_path} (signature={signature_hits}, scale={scale_hits})\n"
    )

    patched_module = types.ModuleType(mixed_module.__name__)
    patched_module.__file__ = str(patched_source_path)
    patched_module.__package__ = mixed_module.__package__
    patched_module.__spec__ = None
    exec(
        compile(patched_source, str(patched_source_path), "exec"),
        patched_module.__dict__,
    )

    sys.modules[mixed_module.__name__] = patched_module
    setattr(kernels_pkg, "mixed_moe_gemm_2stage", patched_module)

    patched_compile_moe_gemm1 = patched_module.compile_mixed_moe_gemm1
    original_compile_flydsl_moe_stage1 = moe_kernels_module.compile_flydsl_moe_stage1

    def patched_compile_flydsl_moe_stage1(
        model_dim,
        inter_dim,
        experts,
        topk,
        tile_m,
        tile_n,
        tile_k,
        doweight_stage1,
        a_dtype,
        b_dtype,
        out_dtype,
    ):
        if b_dtype == "fp4":
            sys.stderr.write(
                "[submission_v88] using file-backed patched compile_mixed_moe_gemm1 "
                f"for tile={tile_m}x{tile_n}x{tile_k}\n"
            )
            return patched_compile_moe_gemm1(
                model_dim=model_dim,
                inter_dim=inter_dim,
                experts=experts,
                topk=topk,
                tile_m=tile_m,
                tile_n=tile_n,
                tile_k=tile_k,
                doweight_stage1=doweight_stage1,
                a_dtype=a_dtype,
                b_dtype=b_dtype,
                out_dtype=out_dtype,
                use_cshuffle_epilog=(out_dtype == "fp8"),
            )
        return original_compile_flydsl_moe_stage1(
            model_dim=model_dim,
            inter_dim=inter_dim,
            experts=experts,
            topk=topk,
            tile_m=tile_m,
            tile_n=tile_n,
            tile_k=tile_k,
            doweight_stage1=doweight_stage1,
            a_dtype=a_dtype,
            b_dtype=b_dtype,
            out_dtype=out_dtype,
        )

    moe_kernels_module.compile_flydsl_moe_stage1 = patched_compile_flydsl_moe_stage1
    patched_compile_moe_gemm1.cache_clear()
    moe_kernels_module._get_compiled_stage1.cache_clear()
    _RUNTIME_FLYDSL_READY = True


def _flydsl_stage1_wrapper(
    hidden_states,
    w1,
    w2,
    sorted_token_ids,
    sorted_expert_ids,
    num_valid_ids,
    out,
    topk,
    kernelName="",
    w1_scale=None,
    a1_scale=None,
    sorted_weights=None,
    **_kwargs,
):
    import aiter

    del w2
    parsed = aiter.ops.flydsl.moe_kernels.get_flydsl_kernel_params(kernelName)
    if parsed is None:
        match = re.fullmatch(
            r"flydsl_moe1_a(?P<a_dtype>[^_]+)_w(?P<b_dtype>[^_]+)_(?P<out_dtype>[^_]+)_t(?P<tile_m>\d+)x(?P<tile_n>\d+)x(?P<tile_k>\d+)",
            kernelName,
        )
        if match is None:
            raise ValueError(f"Invalid FlyDSL kernel name: {kernelName}")
        parsed = {
            "stage": 1,
            "a_dtype": match.group("a_dtype"),
            "b_dtype": match.group("b_dtype"),
            "out_dtype": match.group("out_dtype"),
            "tile_m": int(match.group("tile_m")),
            "tile_n": int(match.group("tile_n")),
            "tile_k": int(match.group("tile_k")),
        }
    return aiter.ops.flydsl.flydsl_moe_stage1(
        a=hidden_states,
        w1=w1,
        sorted_token_ids=sorted_token_ids,
        sorted_expert_ids=sorted_expert_ids,
        num_valid_ids=num_valid_ids,
        out=out,
        topk=topk,
        tile_m=parsed["tile_m"],
        tile_n=parsed["tile_n"],
        tile_k=parsed["tile_k"],
        a_dtype=parsed["a_dtype"],
        b_dtype=parsed["b_dtype"],
        out_dtype=parsed["out_dtype"],
        w1_scale=w1_scale,
        a1_scale=a1_scale,
        sorted_weights=sorted_weights,
    )


def _ensure_flydsl_stage1_patch() -> None:
    global _PATCH_READY
    if _PATCH_READY:
        return

    _ensure_flydsl_runtime_patch()
    fused_moe_module = importlib.import_module("aiter.fused_moe")
    orig_get_2stage_cfgs = fused_moe_module.get_2stage_cfgs

    def patched_get_2stage_cfgs(
        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=True,
    ):
        metadata = copy.copy(
            orig_get_2stage_cfgs(
                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,
            )
        )
        kernel_name = _FLYDSL_STAGE1_BY_TOKEN.get(token)
        if (
            kernel_name is not None
            and expert == 257
            and model_dim == 7168
            and inter_dim == 256
            and topk == 9
        ):
            metadata.stage1 = functools.partial(
                _flydsl_stage1_wrapper,
                kernelName=kernel_name,
            )
        return metadata

    fused_moe_module.get_2stage_cfgs = patched_get_2stage_cfgs
    _PATCH_READY = True


def _use_opus_sorting_for_shape(config: dict) -> bool:
    if (
        config["n_routed_experts"] == 256
        and config["n_shared_experts"] == 1
    ):
        return config["bs"] < 128
    if (
        config["n_routed_experts"] == 32
        and config["n_shared_experts"] == 1
    ):
        if (
            config["d_hidden"] == 7168
            and config["d_expert"] == 512
            and config["n_experts_per_token"] == 8
            and config["bs"] == 16
        ):
            return False
        return config["bs"] < 128
    if (
        config["n_routed_experts"] == 64
        and config["n_shared_experts"] == 1
    ):
        return config["bs"] <= 128
    return False


from aiter import ActivationType, QuantType


def custom_kernel(data: input_t) -> output_t:
    _ensure_custom_fmoe_config()
    _ensure_flydsl_stage1_patch()

    (
        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
    del gate_up_weight, down_weight, gate_up_weight_scale, down_weight_scale

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

    fused_moe_module = importlib.import_module("aiter.fused_moe")
    fused_moe_module._USE_OPUS_MOE_SORTING = _use_opus_sorting_for_shape(config)

    return fused_moe_module.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,
    )
scrolls · 381 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