Skip to content
KernelIndex
Search⌘K

submission 747308

nataliak.6165 · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:91c1f125ca20cb781dc3685ec1bf80f1ea8ab23e2e2f74ea149d610e2f4e6fed
license declaredunknown
license concludedunknown
authorsnataliak.6165
imported2026-08-15

Techniques

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

fp4a_dtype="fp4", b_dtype="fp4", out_dtype="bf16",

Kernel source

submission-v330.py218 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
"""
v330: v323 + patch the runner's FlyDSL stage1 b_scale alias bug and
try the older-stage1-style 64x256x256 tile for E=33 large shapes.
"""
import os, pathlib, re, shutil
os.environ["PYTORCH_ROCM_ARCH"] = "gfx950"
os.environ["AITER_BYPASS_TUNE_CONFIG"] = "1"

# Patch the older runner's mixed stage1 source before importing aiter.
try:
    _src = pathlib.Path("/home/runner/aiter/aiter/ops/flydsl/kernels/mixed_moe_gemm_2stage.py")
    _code = _src.read_text()
    _orig = _code

    if "def compute_f8f6f4_tile(" in _code:
        _code = re.sub(
            r"(a_scale=None,\s*)(b_scale_gate=None,)",
            r"\1b_scale=None,\n                    \2",
            _code,
            count=1,
        )
        _code = re.sub(
            r"(\n\s*)gate_list = list\(acc_gate_in\)",
            (
                r"\1if b_scale is not None:\n"
                r"\1    if b_scale_gate is None:\n"
                r"\1        b_scale_gate = b_scale\n"
                r"\1    if b_scale_up is None:\n"
                r"\1        b_scale_up = b_scale\n"
                r"\1gate_list = list(acc_gate_in)"
            ),
            _code,
            count=1,
        )

    if _code != _orig:
        _src.write_text(_code)
        _pc = _src.parent / "__pycache__"
        if _pc.exists():
            shutil.rmtree(_pc)
        print("Patched FlyDSL stage1 b_scale alias bug", flush=True)
    else:
        print("No FlyDSL stage1 patch applied", flush=True)
except Exception as e:
    print(f"FlyDSL patch error: {e}", flush=True)

import functools, torch

from task import input_t, output_t
from aiter import ActivationType, QuantType

import aiter.fused_moe as _fm
import aiter

from aiter.ops.flydsl.moe_kernels import compile_flydsl_moe_stage2
from aiter.ops.flydsl import flydsl_moe_stage2

HAS_FLYDSL_S1 = False
try:
    from aiter.ops.flydsl.moe_kernels import compile_flydsl_moe_stage1
    from aiter.ops.flydsl import flydsl_moe_stage1

    compile_flydsl_moe_stage1(
        model_dim=7168, inter_dim=512, experts=33, topk=9,
        tile_m=64, tile_n=256, tile_k=256,
        doweight_stage1=False,
        a_dtype="fp4", b_dtype="fp4", out_dtype="bf16",
    )
    compile_flydsl_moe_stage1(
        model_dim=7168, inter_dim=2048, experts=33, topk=9,
        tile_m=64, tile_n=256, tile_k=256,
        doweight_stage1=False,
        a_dtype="fp4", b_dtype="fp4", out_dtype="bf16",
    )
    HAS_FLYDSL_S1 = True
    print("FlyDSL stage1 compiled successfully", flush=True)
except Exception as e:
    print(f"FlyDSL stage1 compile failed: {e}", flush=True)
    import traceback
    traceback.print_exc()
    HAS_FLYDSL_S1 = False

compile_flydsl_moe_stage2(
    model_dim=7168, inter_dim=512, experts=33, topk=9,
    tile_m=64, tile_n=256, tile_k=256,
    doweight_stage2=True, a_dtype="fp4", b_dtype="fp4", out_dtype="bf16",
    accumulate=True,
)
compile_flydsl_moe_stage2(
    model_dim=7168, inter_dim=2048, experts=33, topk=9,
    tile_m=64, tile_n=64, tile_k=256,
    doweight_stage2=True, a_dtype="fp4", b_dtype="fp4", out_dtype="bf16",
    accumulate=True,
)
compile_flydsl_moe_stage2(
    model_dim=7168, inter_dim=256, experts=257, topk=9,
    tile_m=64, tile_n=256, tile_k=256,
    doweight_stage2=True, a_dtype="fp4", b_dtype="fp4", out_dtype="bf16",
    accumulate=True,
)
print("FlyDSL pre-compilation done", flush=True)


def _make_flydsl_stage2(tile_m, tile_n, tile_k=256, mode="reduce"):
    def _flydsl_s2(inter_states, w1, w2, sorted_token_ids, sorted_expert_ids,
                    num_valid_ids, out, topk, kernelName=None, w2_scale=None,
                    a2_scale=None, block_m=32, sorted_weights=None, **kwargs):
        flydsl_moe_stage2(
            inter_states=inter_states, w2=w2,
            sorted_token_ids=sorted_token_ids,
            sorted_expert_ids=sorted_expert_ids,
            num_valid_ids=num_valid_ids,
            out=out, topk=topk,
            tile_m=tile_m, tile_n=tile_n, tile_k=tile_k,
            a_dtype="fp4", b_dtype="fp4", out_dtype="bf16",
            mode=mode,
            w2_scale=w2_scale, a2_scale=a2_scale,
            sorted_weights=sorted_weights,
        )
    return _flydsl_s2


def _make_flydsl_stage1(tile_m, tile_n, tile_k=256):
    def _flydsl_s1(hidden_states, w1, w2, sorted_token_ids, sorted_expert_ids,
                    num_valid_ids, out, topk, kernelName=None, w1_scale=None,
                    a1_scale=None, block_m=32, sorted_weights=None, **kwargs):
        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=tile_m, tile_n=tile_n, tile_k=tile_k,
            a_dtype="fp4", b_dtype="fp4", out_dtype="bf16",
            w1_scale=w1_scale, a1_scale=a1_scale,
            sorted_weights=sorted_weights,
        )
    return _flydsl_s1


_flydsl_s2_S5 = _make_flydsl_stage2(64, 256, 256, "reduce")
_flydsl_s2_S6 = _make_flydsl_stage2(64, 64, 256, "reduce")

if HAS_FLYDSL_S1:
    _flydsl_s1_S5 = _make_flydsl_stage1(64, 256, 256)
    _flydsl_s1_S6 = _make_flydsl_stage1(64, 256, 256)


@functools.lru_cache(maxsize=2048)
def _custom_get_ksplit(token, topk, expert, inter_dim, model_dim):
    if token <= 128:
        return 2
    return 0


_fm.get_ksplit = _custom_get_ksplit
_orig_get_block_size_M = getattr(_fm.get_block_size_M, "__wrapped__", _fm.get_block_size_M)


@functools.lru_cache(maxsize=2048)
def _custom_get_block_size_M(token, topk, expert, inter_dim):
    if token <= 128:
        return 16
    if expert > 128:
        return 32
    if inter_dim > 1024:
        return 128
    return _orig_get_block_size_M(token, topk, expert, inter_dim)


_fm.get_block_size_M = _custom_get_block_size_M
_orig_get_2stage_cfgs = _fm.get_2stage_cfgs


def _patched_get_2stage_cfgs(*args, **kwargs):
    cfg = _orig_get_2stage_cfgs(*args, **kwargs)
    if cfg is None:
        return cfg

    if cfg.ksplit == 0:
        inter_dim = args[2]
        expert = args[3]

        if inter_dim > 1024:
            cfg.stage2 = _flydsl_s2_S6
        else:
            cfg.stage2 = _flydsl_s2_S5

        if HAS_FLYDSL_S1 and expert <= 128:
            if inter_dim > 1024:
                cfg.stage1 = _flydsl_s1_S6
            else:
                cfg.stage1 = _flydsl_s1_S5

    return cfg


_fm.get_2stage_cfgs = _patched_get_2stage_cfgs

from aiter.fused_moe import fused_moe


def custom_kernel(data: input_t) -> output_t:
    (hs, rw1, rw2, rw1s, rw2s, w1, w2, w1s, w2s, tw, ti, cfg) = data

    hp = cfg.get("d_hidden_pad", cfg["d_hidden"]) - cfg["d_hidden"]
    ip = cfg.get("d_expert_pad", cfg["d_expert"]) - cfg["d_expert"]

    return fused_moe(
        hs, w1, w2, tw, ti,
        expert_mask=None, activation=ActivationType.Silu,
        quant_type=QuantType.per_1x32, doweight_stage1=False,
        w1_scale=w1s, w2_scale=w2s, a1_scale=None, a2_scale=None,
        hidden_pad=hp, intermediate_pad=ip,
    )
scrolls · 218 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