Skip to content
KernelIndex
Search⌘K

submission 581053

RyanWillie · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:8901a1d2672abfef8db321a645a3a9e710c3ddd0e2ec3ae417580922882740dc
license declaredunknown
license concludedunknown
authorsRyanWillie
imported2026-08-15

Techniques

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

fp4DeepSeek-R1 style MXFP4 MoE: v006 selective ksplit + FlyDSL stage2 (tile_m=32) + JIT warmup.

Kernel source

submission.py272 lines
# Zeus Variant: v037_flydsl_tile_m32_warmup
# Experiment: #037
# Technique: FlyDSL stage2 tile_m=32 for bs512 shapes + JIT warmup at import time
# Mechanism: Same as v034 (tile_m=32) but pre-compiles FlyDSL MLIR kernels at import
#            time so compilation occurs during test phase. By benchmark/ranked phases
#            the kernels are cached in @functools.cache and execute instantly.
#            Small shapes use v006 selective ksplit=2 (CKTile path, no stage2 change needed).
#            bs512/e256 uses CSV-tuned CK kernel, left unchanged.

"""
DeepSeek-R1 style MXFP4 MoE: v006 selective ksplit + FlyDSL stage2 (tile_m=32) + JIT warmup.
"""

#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X

import os
import functools
from contextlib import contextmanager
from dataclasses import dataclass
from typing import Callable

from task import input_t, output_t

from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import fused_moe, get_ksplit
import aiter.fused_moe as _fused_moe_mod
import aiter.ops.flydsl as _flydsl_mod


# ---- JIT warmup: pre-compile FlyDSL MLIR kernels in background thread ----
# FlyDSL MLIR compilation for tile_m=32 is too slow to fit in leaderboard timeout
# if it runs inline. By starting compilation in a background thread at import time,
# it overlaps with CK sorting (~25s) + CKTile (~106s) + CK 2-stage (~104s) JIT
# that happens lazily on first custom_kernel() call. By the time the benchmark
# reaches bs512/e512 or bs512/e2048 shapes, the FlyDSL kernels are already cached
# in _get_compiled_stage2's @functools.cache.

import threading as _threading
import inspect as _inspect

from aiter.ops.flydsl.moe_kernels import _get_compiled_stage2 as _warmup_compile

# Detect server AITER signature: older version lacks `persist_m` parameter
_warmup_sig = _inspect.signature(_warmup_compile.__wrapped__)
_has_persist_m = "persist_m" in _warmup_sig.parameters

# Pre-compile for each bs512 shape that uses FlyDSL stage2.
# Parameters must exactly match what flydsl_moe_stage2() will compute at runtime:
#   model_dim=7168 (d_hidden_pad), topk=9, experts=33
#   a_dtype="fp4" => inter_dim = d_expert_pad (already doubled by flydsl_moe_stage2)
#   doweight=True (sorted_weights are passed), accumulate=True (mode="atomic")
#   persist_m=4 (sorted_expert_ids.numel() > 256 for bs512) -- only if server supports it
_WARMUP_CONFIGS_FULL = [
    # (model_dim, inter_dim, experts, topk, tile_m, tile_n, tile_k,
    #  doweight, a_dtype, b_dtype, out_dtype, accumulate, persist_m)
    (7168, 512, 33, 9, 32, 128, 256, True, "fp4", "fp4", "bf16", True, 4),   # bs512/e512
    (7168, 2048, 33, 9, 32, 128, 256, True, "fp4", "fp4", "bf16", True, 4),  # bs512/e2048
]

# Strip persist_m if server doesn't support it
_WARMUP_CONFIGS = [cfg if _has_persist_m else cfg[:-1] for cfg in _WARMUP_CONFIGS_FULL]

def _do_warmup():
    """Compile FlyDSL kernels. Runs in background thread to overlap with CK JIT."""
    for cfg in _WARMUP_CONFIGS:
        _warmup_compile(*cfg)

_warmup_thread = _threading.Thread(target=_do_warmup, daemon=True)
_warmup_thread.start()

# NOTE: We do NOT join here — let it run in background while CK/CKTile compile.
# We join before the first FlyDSL call in the stage2 wrapper below.


# ---- Build FlyDSL stage2 replacement ----

def _build_flydsl_stage2(tile_m=64, tile_n=128, tile_k=256, mode="atomic"):
    """Build a FlyDSL stage2 callable matching ck_moe_stage2_fwd interface."""
    def _flydsl_stage2(
        inter_states, w1, w2,
        sorted_token_ids, sorted_expert_ids, num_valid_ids,
        out, topk,
        kernelName="", w2_scale=None, a2_scale=None,
        sorted_weights=None, **_kw,
    ):
        # Ensure warmup compilation is done before first FlyDSL call.
        # After first join, _warmup_thread.is_alive() is False and join() is a no-op.
        _warmup_thread.join()
        _flydsl_mod.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_stage2


# ---- Patch get_2stage_cfgs to inject FlyDSL for bs512 shapes ----

_orig_get_2stage_cfgs = _fused_moe_mod.get_2stage_cfgs.__wrapped__

# Shape keys (padded_M, inter_dim, E) to apply FlyDSL stage2
_FLYDSL_TARGETS = {
    (512, 512, 33),   # bs512/e512
    (512, 2048, 33),  # bs512/e2048
}

_flydsl_stage2_fn = _build_flydsl_stage2(tile_m=32, tile_n=128, tile_k=256, mode="atomic")


@functools.lru_cache(maxsize=2048)
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,
):
    orig = _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,
    )

    # Only replace stage2 for targeted bs512 shapes with ksplit=0
    if (token, inter_dim, expert) in _FLYDSL_TARGETS and orig.ksplit == 0:
        # Replace stage2 with FlyDSL, keep everything else
        @dataclass
        class PatchedMetadata:
            stage1: Callable
            stage2: Callable
            block_m: int
            ksplit: int
            run_1stage: bool = False
            has_bias: bool = False
            use_non_temporal_load: bool = True

        return PatchedMetadata(
            stage1=orig.stage1,
            stage2=_flydsl_stage2_fn,
            block_m=orig.block_m,
            ksplit=orig.ksplit,
            run_1stage=orig.run_1stage,
        )

    return orig


_fused_moe_mod.get_2stage_cfgs = _patched_get_2stage_cfgs


# ---- v006 selective ksplit ----

FORCED_KSPLIT = "2"
_ENV_OVERRIDES = {
    "AITER_BYPASS_TUNE_CONFIG": "1",
    "AITER_KSPLIT": FORCED_KSPLIT,
}


@contextmanager
def _temporary_env(overrides: dict[str, str]):
    previous = {key: os.environ.get(key) for key in overrides}
    try:
        for key, value in overrides.items():
            os.environ[key] = value
        yield
    finally:
        for key, value in previous.items():
            if value is None:
                os.environ.pop(key, None)
            else:
                os.environ[key] = value


def _clear_ksplit_cache() -> None:
    get_ksplit.cache_clear()


def _should_force_ksplit(config: dict) -> bool:
    return config["bs"] <= 128 and config["d_expert"] <= 512


def _run_fused_moe(
    hidden_states, gate_up_weight_shuffled, down_weight_shuffled,
    topk_weights, topk_ids,
    gate_up_weight_scale_shuffled, down_weight_scale_shuffled,
    hidden_pad, intermediate_pad,
):
    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,
    )


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

    if not _should_force_ksplit(config):
        return _run_fused_moe(
            hidden_states,
            gate_up_weight_shuffled,
            down_weight_shuffled,
            topk_weights,
            topk_ids,
            gate_up_weight_scale_shuffled,
            down_weight_scale_shuffled,
            hidden_pad,
            intermediate_pad,
        )

    _clear_ksplit_cache()
    try:
        with _temporary_env(_ENV_OVERRIDES):
            _clear_ksplit_cache()
            return _run_fused_moe(
                hidden_states,
                gate_up_weight_shuffled,
                down_weight_shuffled,
                topk_weights,
                topk_ids,
                gate_up_weight_scale_shuffled,
                down_weight_scale_shuffled,
                hidden_pad,
                intermediate_pad,
            )
    finally:
        _clear_ksplit_cache()
scrolls · 272 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