Skip to content
KernelIndex
Search⌘K

submission 707719

internetrat · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

moe-optimized-v48-cktile-blockm-kernels-ep2048-only.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-707719?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
169.5µs
#307 of 782
2026-04-03

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:7df6c984e6656e36deda2145bda6b0106d645b50ebb49f76b594b8cc7412d992
license declaredunknown
license concludedunknown
authorsinternetrat
imported2026-08-15

Kernel source

moe-optimized-v48-cktile-blockm-kernels-ep2048-only.py195 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X

import importlib.util
import os
import pathlib
import re


def ensure_seed_cfg(path: str):
    if os.path.exists(path):
        return
    content = """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,0,40.8894,moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16,0.0%,27.8142,moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16,1.3%,68.7036,0,23.08,20597.69,
256,128,7168,256,257,9,ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,32,0,87.2657,moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16,0.0%,52.5931,moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16,1.3%,139.8588,0,90.69,10135.53,
256,512,7168,256,257,9,ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,32,0,97.693,moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16,0.0%,70.2081,moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16,1.3%,167.9011,0,302.17,8491.91,
256,16,7168,512,33,9,ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,32,0,0.0,moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16,0.0%,0.0,moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16,2.9%,47.3341,0,66.99,7683.16,
256,128,7168,512,33,9,ActivationType.Silu,torch.bfloat16,torch.float4_e2m1fn_x2,torch.float4_e2m1fn_x2,QuantType.per_1x32,1,0,32,0,0.0,moe_ck2stages_gemm1_64x32x32x128_1x1_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16,0.0%,0.0,moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16,2.9%,58.2354,0,435.6,6286.28,
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,0.0,moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16,0.0%,0.0,moe_ck2stages_gemm2_64x32x32x128_1x1_MulABScaleExpertWeightShuffled_v1_Nswizzle0_Quant3_MulRoutedWeight1_FP4X2_FP4X2_B16,2.9%,128.7392,0,788.17,2907.75,
"""
    with open(path, "w", encoding="utf-8") as f:
        f.write(content)


def _patch_file(path: str, transform) -> bool:
    if not os.path.exists(path):
        return False
    txt = pathlib.Path(path).read_text(encoding="utf-8")
    txt2 = transform(txt)
    if txt2 == txt:
        return False
    pathlib.Path(path).write_text(txt2, encoding="utf-8")
    return True


def patch_runtime():
    from aiter.jit.core import AITER_CSRC_DIR

    ck_dir = os.path.join(AITER_CSRC_DIR, "ck_gemm_moe_2stages_codegen")
    ck_common_p = os.path.join(ck_dir, "gemm_moe_ck2stages_common.py")
    ck_tune_p = os.path.join(ck_dir, "gemm_moe_tune.py")

    _patch_file(
        ck_common_p,
        lambda t: t.replace(
            "#  3: kernelInstanceGEMM1(       256,      256,         128,       128,     2,       2,        3,),",
            "   3: kernelInstanceGEMM1(       256,      256,         128,       128,     2,       2,        3,),",
        ).replace(
            "#  7: kernelInstanceGEMM2(      256,       256,         64,       128,     2,       2,         3,),",
            "   7: kernelInstanceGEMM2(      256,       256,         64,       128,     2,       2,         3,),",
        ),
    )

    _patch_file(
        ck_tune_p,
        lambda t: (
            t.replace(
                "blockMs = [16, 32, 64, 128]",
                "blockMs = [16, 32, 64, 128, 256] if get_gfx() in [\"gfx950\"] else [16, 32, 64, 128]",
            )
            .replace(
                "if blockM in [16, 32, 64, 128] and use_g1u1:",
                "if blockM in [16, 32, 64, 128, 256] and use_g1u1:",
            )
            .replace(
                "if blockM not in [16, 32, 64, 128] or not use_g1u1:",
                "if blockM not in [16, 32, 64, 128, 256] or not use_g1u1:",
            )
        ),
    )

    cktile_dir = os.path.join(AITER_CSRC_DIR, "ck_tile_gemm_moe_2stages")
    cktile_common_p = os.path.join(cktile_dir, "moe_cktile2stages_common.py")

    def _patch_cktile_common(t: str) -> str:
        s1_anchor = "3: kernelInstance(       1,        256,       64,        256,       256,           16,         16,          32,          1,           4,          1,),"
        if s1_anchor in t and "kernelInstance(       1,        256,      256,        256,       256," not in t:
            t = t.replace(
                s1_anchor,
                "\n".join(
                    [
                        s1_anchor,
                        "    4: kernelInstance(       1,        256,      128,        256,       256,           16,         16,          32,          1,           4,          1,),",
                        "    5: kernelInstance(       1,        256,      256,        256,       256,           16,         16,          32,          1,           4,          1,),",
                    ]
                ),
            )

        s2_anchor = "3: kernelInstance(       2,        256,       64,        256,       256,           16,         16,          32,          1,        4,            1,),"
        if s2_anchor in t and "kernelInstance(       2,        256,      256,        256,       256," not in t:
            t = t.replace(
                s2_anchor,
                "\n".join(
                    [
                        s2_anchor,
                        "    4: kernelInstance(       2,        256,      128,        256,       256,           16,         16,          32,          1,        4,            1,),",
                        "    5: kernelInstance(       2,        256,      256,        256,       256,           16,         16,          32,          1,        4,            1,),",
                    ]
                ),
            )
        return t

    _patch_file(cktile_common_p, _patch_cktile_common)

    spec = importlib.util.find_spec("aiter.fused_moe")
    if spec is None or spec.origin is None:
        return
    fused_moe_p = spec.origin

    def _patch_fused_moe_py(t: str) -> str:
        if "def get_cktile_block_m" in t:
            return t

        t = re.sub(
            r"def get_block_m\(\) -> int:\n(?:\s+.*\n)+?\s+return 16 if token < 2048 else 32 if token < 16384 else 64\n",
            "def get_cktile_block_m() -> int:\n"
            "        if cfg is not None:\n"
            "            bm = block_m\n"
            "        else:\n"
            "            if q_dtype_a == dtypes.fp8:\n"
            "                bm = 32\n"
            "            else:\n"
            "                bm = 16 if token < 2048 else 32 if token < 16384 else 64\n"
            "        allowed = [16, 32, 64, 128, 256] if get_gfx() == \"gfx950\" else [16, 32, 64]\n"
            "        if bm in allowed:\n"
            "            return bm\n"
            "        smaller = [m for m in allowed if m <= bm]\n"
            "        return smaller[-1] if smaller else allowed[0]\n",
            t,
            count=1,
            flags=re.MULTILINE,
        )

        t = t.replace("get_block_m(),\n            ksplit,", "get_cktile_block_m(),\n            ksplit,")
        t = t.replace(
            "16 if token < 2048 else 32 if token < 16384 else 64,\n            ksplit,",
            "get_cktile_block_m(),\n            ksplit,",
        )
        return t

    _patch_file(fused_moe_p, _patch_fused_moe_py)


seed_cfg = "/tmp/seed_fmoe_no_ep2048.csv"
ensure_seed_cfg(seed_cfg)

os.environ["AITER_CONFIG_FMOE"] = seed_cfg
os.environ["AITER_ONLINE_TUNE"] = "1"

from task import input_t, output_t

patch_runtime()

from aiter import ActivationType, QuantType
from aiter.fused_moe import fused_moe


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"]

    return 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 · 195 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