Skip to content
KernelIndex
Search⌘K

submission 754350

NinoHeather · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

my_submission_46.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-754350?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
121.5µs
#44 of 782
2026-04-07

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:a876e76d3d093db5f81e387e0fff7b74994a67cde2c388a58163d785d5cb1957
license declaredunknown
license concludedunknown
authorsNinoHeather
imported2026-08-15

Techniques

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

fp4"a_dtype": "fp4",

Kernel source

my_submission_46.py704 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
from __future__ import annotations

import functools
import inspect
import math
import os
from dataclasses import dataclass
from typing import Any, Dict, Iterable, Tuple

import pandas as pd
import torch

from task import input_t, output_t

import aiter
import aiter.fused_moe as fmoe
import aiter.ops.flydsl.moe_kernels as flydsl_kernels
from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import fused_moe
from aiter.ops.moe_sorting import moe_sorting_fwd
from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort


torch.set_grad_enabled(False)


def _apply_runtime_env() -> None:
    defaults = {
        "PYTORCH_ROCM_ARCH": "gfx950",
        "AITER_USE_NT": "1",
        "AITER_USE_FLYDSL_MOE": "1",
        "AITER_USE_FLYDSL_MOE_STAGE2": "1",
        "AITER_USE_OPUS_MOE_SORTING": "1",
    }
    for k, v in defaults.items():
        os.environ.setdefault(k, v)


def _register_missing_flydsl_kernels() -> None:
    flydsl_kernels._KERNEL_PARAMS.setdefault(
        "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic",
        {
            "stage": 2,
            "a_dtype": "fp4",
            "b_dtype": "fp4",
            "out_dtype": "bf16",
            "tile_m": 16,
            "tile_n": 128,
            "tile_k": 128,
            "mode": "atomic",
            "MPerBlock": 16,
        },
    )


_SIG_FIELDS = (
    "cu_num",
    "token",
    "model_dim",
    "inter_dim",
    "expert",
    "topk",
    "act_type",
    "dtype",
    "q_dtype_a",
    "q_dtype_w",
    "q_type",
    "use_g1u1",
    "doweight_stage1",
)
_SIG_META = (
    "ActivationType.Silu",
    "torch.bfloat16",
    "torch.float4_e2m1fn_x2",
    "torch.float4_e2m1fn_x2",
    "QuantType.per_1x32",
)

_K1_256x128 = (
    "moe_ck2stages_gemm1_256x128x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_"
    "Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
_K1_256x32 = (
    "moe_ck2stages_gemm1_256x32x128x128_1x4_MulABScaleShuffled_v3_Nswizzle0_"
    "Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
_K2_F16_128_A = "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic"
_K2_FLY_64_REDUCE = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce"


@dataclass(frozen=True)
class ShapePlan:
    token: int
    model_dim: int
    inter_dim: int
    expert: int
    topk: int = 9
    block_m: int = 32
    ksplit: int = 0
    kernel1: str = ""
    kernel2: str = ""
    run_1stage: bool = False
    use_nt: bool | None = None


def _signature(plan: ShapePlan) -> Tuple[Any, ...]:
    return (
        256,
        plan.token,
        plan.model_dim,
        plan.inter_dim,
        plan.expert,
        plan.topk,
        *_SIG_META,
        1,
        0,
    )


def _row_data(plan: ShapePlan) -> Dict[str, Any]:
    row = {
        "block_m": plan.block_m,
        "ksplit": plan.ksplit,
        "kernelName1": plan.kernel1,
        "kernelName2": plan.kernel2,
        "run_1stage": plan.run_1stage,
        "us1": 0.0,
        "err1": "0.0%",
        "us2": 0.0,
        "err2": "0.0%",
        "us": -1.0,
        "tflops": 0.0,
        "bw": 0.0,
        "_tag": "",
    }
    if plan.use_nt is not None:
        row["use_non_temporal_load"] = bool(plan.use_nt)
    return row


def _retarget_stage_kernel(stage: Any, *, kernel_name: str, block_m: int) -> Any:
    if not isinstance(stage, functools.partial):
        return stage
    try:
        params = set(inspect.signature(stage.func).parameters)
    except Exception:
        params = set()
    kw = dict(stage.keywords or {})
    if "kernelName" in kw or "kernelName" in params or not params:
        kw["kernelName"] = kernel_name
    if "block_m" in kw or "block_m" in params or not params:
        kw["block_m"] = block_m
    return functools.partial(stage.func, *stage.args, **kw)


def _supports_kernel_retarget(stage: Any) -> bool:
    if not isinstance(stage, functools.partial):
        return False
    try:
        params = set(inspect.signature(stage.func).parameters)
    except Exception:
        params = set()
    kw = dict(stage.keywords or {})
    return ("kernelName" in params) or ("kernelName" in kw) or (not params)


def _manual_plans() -> Iterable[ShapePlan]:
    model_dim = 7168
    yield ShapePlan(16, model_dim, 256, 257, block_m=16, ksplit=2)
    yield ShapePlan(128, model_dim, 256, 257, block_m=16, ksplit=2)
    yield ShapePlan(
        512,
        model_dim,
        256,
        257,
        block_m=32,
        ksplit=0,
        kernel1=_K1_256x32,
        kernel2=_K2_F16_128_A,
        use_nt=True,
    )
    yield ShapePlan(16, model_dim, 512, 33, block_m=32, ksplit=2)
    yield ShapePlan(128, model_dim, 512, 33, block_m=64, ksplit=0, kernel1=_K1_256x128, kernel2=_K2_F16_128_A)
    yield ShapePlan(512, model_dim, 512, 33, block_m=64, ksplit=0, kernel1=_K1_256x128, kernel2=_K2_F16_128_A)
    yield ShapePlan(512, model_dim, 2048, 33, block_m=64, ksplit=0, kernel1=_K1_256x128, kernel2=_K2_F16_128_A)


def _read_csv_configs() -> Dict[Tuple[Any, ...], Dict[str, Any]]:
    cfg_root = os.path.join(os.path.dirname(aiter.__file__), "configs")
    model_cfg_root = os.path.join(cfg_root, "model_configs")
    csv_files = [os.path.join(cfg_root, "tuned_fmoe.csv")]
    if os.path.isdir(model_cfg_root):
        for name in sorted(os.listdir(model_cfg_root)):
            if name.endswith(".csv") and "fmoe" in name:
                csv_files.append(os.path.join(model_cfg_root, name))

    frames = [pd.read_csv(path) for path in csv_files if os.path.exists(path)]
    if not frames:
        existing = getattr(fmoe, "cfg_2stages", None)
        return dict(existing) if isinstance(existing, dict) else {}

    merged = pd.concat(frames, ignore_index=True)
    merged = merged.drop_duplicates(subset=list(_SIG_FIELDS), keep="last")
    for col in (
        "ksplit",
        "block_m",
        "cu_num",
        "token",
        "model_dim",
        "inter_dim",
        "expert",
        "topk",
        "use_g1u1",
        "doweight_stage1",
        "run_1stage",
    ):
        if col in merged.columns:
            merged[col] = merged[col].fillna(0).astype(int)

    util = (merged["token"] * merged["topk"]) // merged["expert"].clip(lower=1)
    merged.loc[(merged["expert"] > 0) & (merged["topk"] > 0) & (util <= 10), "ksplit"] = 4

    kill = (merged["expert"] == 257) & (merged["inter_dim"] == 256) & (merged["token"].isin([16, 128]))
    merged = merged[~kill]

    cfg = merged.set_index(list(_SIG_FIELDS)).to_dict("index")
    for row in cfg.values():
        for k, v in list(row.items()):
            if isinstance(v, float) and math.isnan(v):
                row[k] = ""
            elif isinstance(v, float) and v == int(v):
                row[k] = int(v)
    return cfg


def _apply_plans_into_cfg(cfg: Dict[Tuple[Any, ...], Dict[str, Any]]) -> None:
    for plan in _manual_plans():
        cfg[_signature(plan)] = _row_data(plan)


def _patch_scheduler_heuristics() -> None:
    fmoe.use_nt = lambda token, topk, e: True  # type: ignore[assignment]

    def ksplit_rule(token: int, topk: int, expert: int, inter_dim: int, model_dim: int) -> int:
        if expert == 33 and token >= 512:
            return 0
        dense = (token * topk) // max(expert, 1)
        if dense <= 10:
            return 4
        if dense <= 64:
            return 2
        return 0

    @functools.lru_cache(maxsize=1024)
    def block_m_rule(token: int, topk: int, expert: int, inter_dim: int) -> int:
        dense = (token * topk) // max(expert, 1)
        if dense >= 100:
            return 128
        if dense >= 20:
            return 64
        return 32

    fmoe.get_ksplit = ksplit_rule  # type: ignore[assignment]
    fmoe.get_block_size_M = block_m_rule  # type: ignore[assignment]


def _patch_runtime_nt_from_cfg() -> None:
    base_get_cfg = fmoe.get_2stage_cfgs
    enable_expert257_bs128_tuned = True

    @functools.lru_cache(maxsize=2048)
    def wrapped(
        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,
    ):
        meta = base_get_cfg(
            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,
        )
        try:
            from aiter.jit.utils.chip_info import get_cu_num
        except Exception:
            return meta

        sig = (
            get_cu_num(),
            token,
            model_dim,
            inter_dim,
            expert,
            topk,
            str(activation),
            str(dtype),
            str(q_dtype_a),
            str(q_dtype_w),
            str(q_type),
            use_g1u1,
            doweight_stage1,
        )
        row = fmoe.cfg_2stages.get(sig) if isinstance(fmoe.cfg_2stages, dict) else None
        out = meta
        if row and row.get("use_non_temporal_load") is not None:
            flag = bool(row["use_non_temporal_load"])

            def maybe_toggle(stage):
                if not isinstance(stage, functools.partial):
                    return stage
                kw = dict(stage.keywords or {})
                if "use_non_temporal_load" in kw:
                    kw["use_non_temporal_load"] = flag
                    return functools.partial(stage.func, *stage.args, **kw)
                if "non_temporal_load" in kw:
                    kw["non_temporal_load"] = flag
                    return functools.partial(stage.func, *stage.args, **kw)
                return stage

            s1 = maybe_toggle(out.stage1)
            s2 = maybe_toggle(out.stage2)
            out = fmoe.MOEMetadata(s1, s2, out.block_m, out.ksplit, out.run_1stage, out.has_bias, flag)

        if (
            enable_expert257_bs128_tuned
            and token == 128
            and model_dim == 7168
            and inter_dim == 256
            and expert == 257
            and topk == 9
            and _supports_kernel_retarget(out.stage1)
            and _supports_kernel_retarget(out.stage2)
        ):
            s1 = _retarget_stage_kernel(out.stage1, kernel_name=_K1_256x128, block_m=32)
            s2 = _retarget_stage_kernel(out.stage2, kernel_name=_K2_FLY_64_REDUCE, block_m=32)
            out = fmoe.MOEMetadata(s1, s2, 32, 1, False, out.has_bias, False)

        return out

    fmoe.get_2stage_cfgs = wrapped  # type: ignore[assignment]


def _bootstrap() -> None:
    _apply_runtime_env()
    _register_missing_flydsl_kernels()
    _patch_scheduler_heuristics()
    cfg = _read_csv_configs()
    _apply_plans_into_cfg(cfg)
    fmoe.cfg_2stages = cfg
    if hasattr(fmoe.get_2stage_cfgs, "cache_clear"):
        fmoe.get_2stage_cfgs.cache_clear()
    _patch_runtime_nt_from_cfg()
    original_get_padded_m = getattr(fmoe, "get_padded_M", None)
    if original_get_padded_m is not None:
        fmoe.get_padded_M = functools.lru_cache(maxsize=256)(original_get_padded_m)  # type: ignore[assignment]


def _install_expert33_dp2_sorting() -> None:
    orig_impl = getattr(fmoe, "_moe_sorting_impl", None)
    sort_fn = getattr(aiter, "moe_sorting_opus_fwd", None)
    if orig_impl is None or sort_fn is None:
        return

    workspace_cache: Dict[Tuple[Any, ...], Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]] = {}

    def _wrapped(
        topk_ids,
        topk_weights,
        num_experts,
        model_dim,
        moebuf_dtype,
        block_size,
        expert_mask=None,
        num_local_tokens=None,
        dispatch_policy=0,
        use_opus=False,
    ):
        if int(num_experts) != 33:
            return orig_impl(
                topk_ids,
                topk_weights,
                num_experts,
                model_dim,
                moebuf_dtype,
                block_size,
                expert_mask,
                num_local_tokens,
                dispatch_policy,
                use_opus,
            )

        m_tokens, topk = topk_ids.shape
        max_padded = int(topk_ids.numel() + num_experts * block_size - topk)
        max_blocks = int((max_padded + block_size - 1) // block_size)
        device = topk_ids.device
        key = (
            device.index if device.type == "cuda" else -1,
            max_padded,
            max_blocks,
            int(m_tokens),
            int(model_dim),
            str(moebuf_dtype),
            int(block_size),
            int(num_experts),
        )
        cached = workspace_cache.get(key)
        if cached is None:
            cached = (
                torch.empty(max_padded, dtype=torch.int32, device=device),
                torch.empty(max_padded, dtype=torch.float32, device=device),
                torch.empty(max_blocks, dtype=torch.int32, device=device),
                torch.empty(2, dtype=torch.int32, device=device),
                torch.empty((m_tokens, model_dim), dtype=moebuf_dtype, device=device),
            )
            workspace_cache[key] = cached

        sid, sw, sei, nvi, mb = cached
        sort_fn(
            topk_ids,
            topk_weights,
            sid,
            sw,
            sei,
            nvi,
            mb,
            num_experts,
            block_size,
            expert_mask,
            num_local_tokens,
            2,
        )
        return sid, sw, sei, nvi, mb

    fmoe._moe_sorting_impl = _wrapped


_bootstrap()
_install_expert33_dp2_sorting()


@dataclass
class _DirectBS512Expert33Workspace:
    meta: Any
    block_m: int
    token: int
    topk: int
    inter_dim: int
    total_experts: int
    gate_up_shuffled: torch.Tensor
    down_shuffled: torch.Tensor
    w1_scale: torch.Tensor
    w2_scale: torch.Tensor
    sorted_ids: torch.Tensor
    sorted_weights: torch.Tensor
    sorted_expert_ids: torch.Tensor
    num_valid_ids: torch.Tensor
    moe_out: torch.Tensor
    a2_buf: torch.Tensor


_DIRECT_BS512_EXPERT33_REGISTRY: Dict[Tuple[Any, ...], _DirectBS512Expert33Workspace | None] = {}


def _is_bs512_expert33_target(cfg: Dict[str, Any], token: int, expert: int, inter_dim: int) -> bool:
    return (
        token == 512
        and expert == 33
        and inter_dim in (512, 2048)
        and int(cfg.get("n_routed_experts", 0)) == 32
        and int(cfg.get("n_shared_experts", 0)) == 1
        and int(cfg.get("d_hidden", 0)) == 7168
    )


def _direct_registry_key(data: input_t) -> Tuple[Any, ...]:
    hs, w1, w2, s1, s2, cfg = data[0], data[5], data[6], data[7], data[8], data[11]
    dev = hs.device
    return (
        dev.type,
        dev.index if dev.type == "cuda" else -1,
        int(cfg["d_expert"]),
        int(w1.data_ptr()),
        int(w2.data_ptr()),
        int(s1.data_ptr()),
        int(s2.data_ptr()),
    )


def _as_fp8_e8m0_scale(weight: torch.Tensor, scale: torch.Tensor) -> torch.Tensor:
    return scale.view(dtypes.fp8_e8m0) if weight.dtype == dtypes.fp4x2 else scale


def _make_direct_workspace(data: input_t) -> _DirectBS512Expert33Workspace | None:
    hs, w1, w2, s1, s2, tw, cfg = data[0], data[5], data[6], data[7], data[8], data[9], data[11]
    token = int(hs.shape[0])
    total_experts = int(w1.shape[0])
    topk = int(tw.shape[1])
    hidden_pad = int(cfg["d_hidden_pad"] - cfg["d_hidden"])
    inter_pad = int(cfg["d_expert_pad"] - cfg["d_expert"])
    _, model_dim, inter_dim = fmoe.get_inter_dim(w1.shape, w2.shape)

    meta = fmoe.get_2stage_cfgs(
        fmoe.get_padded_M(token),
        model_dim,
        inter_dim,
        total_experts,
        topk,
        hs.dtype,
        dtypes.fp4x2,
        w1.dtype,
        QuantType.per_1x32,
        inter_dim != w1.shape[1],
        ActivationType.Silu,
        False,
        hidden_pad,
        inter_pad,
        getattr(w1, "is_shuffled", False),
    )
    if meta.run_1stage:
        return None

    block_m = int(meta.block_m)
    padded_tokens = int(token * topk + total_experts * block_m - topk)
    padded_blocks = int((padded_tokens + block_m - 1) // block_m)
    device = hs.device
    return _DirectBS512Expert33Workspace(
        meta=meta,
        block_m=block_m,
        token=token,
        topk=topk,
        inter_dim=int(inter_dim),
        total_experts=total_experts,
        gate_up_shuffled=w1,
        down_shuffled=w2,
        w1_scale=_as_fp8_e8m0_scale(w1, s1),
        w2_scale=_as_fp8_e8m0_scale(w2, s2),
        sorted_ids=torch.empty(padded_tokens, dtype=torch.int32, device=device),
        sorted_weights=torch.empty(padded_tokens, dtype=torch.float32, device=device),
        sorted_expert_ids=torch.empty(padded_blocks, dtype=torch.int32, device=device),
        num_valid_ids=torch.empty(2, dtype=torch.int32, device=device),
        moe_out=torch.empty((token, int(model_dim)), dtype=hs.dtype, device=device),
        a2_buf=torch.empty((token, topk, int(inter_dim)), dtype=hs.dtype, device=device),
    )


def _execute_direct_workspace(
    ws: _DirectBS512Expert33Workspace,
    hidden_states: torch.Tensor,
    topk_weights: torch.Tensor,
    topk_ids: torch.Tensor,
) -> output_t:
    moe_sorting_fwd(
        topk_ids,
        topk_weights,
        ws.sorted_ids,
        ws.sorted_weights,
        ws.sorted_expert_ids,
        ws.num_valid_ids,
        ws.moe_out,
        ws.total_experts,
        ws.block_m,
        None,
        None,
        0,
    )
    q_hidden, q_hidden_scale = fused_dynamic_mxfp4_quant_moe_sort(
        hidden_states,
        sorted_ids=ws.sorted_ids,
        num_valid_ids=ws.num_valid_ids,
        token_num=ws.token,
        topk=1,
        block_size=ws.block_m,
    )
    ws.meta.stage1(
        q_hidden,
        ws.gate_up_shuffled,
        ws.down_shuffled,
        ws.sorted_ids,
        ws.sorted_expert_ids,
        ws.num_valid_ids,
        ws.a2_buf,
        ws.topk,
        block_m=ws.block_m,
        a1_scale=q_hidden_scale,
        w1_scale=ws.w1_scale,
        sorted_weights=None,
    )
    q_a2, q_a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
        ws.a2_buf.view(-1, ws.inter_dim),
        sorted_ids=ws.sorted_ids,
        num_valid_ids=ws.num_valid_ids,
        token_num=ws.token,
        topk=ws.topk,
        block_size=ws.block_m,
    )
    ws.meta.stage2(
        q_a2.view(ws.token, ws.topk, -1),
        ws.gate_up_shuffled,
        ws.down_shuffled,
        ws.sorted_ids,
        ws.sorted_expert_ids,
        ws.num_valid_ids,
        ws.moe_out,
        ws.topk,
        w2_scale=ws.w2_scale,
        a2_scale=q_a2_scale,
        block_m=ws.block_m,
        sorted_weights=ws.sorted_weights,
    )
    return ws.moe_out


def _try_direct_bs512_expert33(data: input_t) -> output_t | None:
    cache_key = _direct_registry_key(data)
    try:
        cached_workspace = _DIRECT_BS512_EXPERT33_REGISTRY[cache_key]
    except KeyError:
        cached_workspace = _make_direct_workspace(data)
        _DIRECT_BS512_EXPERT33_REGISTRY[cache_key] = cached_workspace

    if cached_workspace is None:
        return None

    hidden_states = data[0]
    routed_weights = data[9]
    routed_ids = data[10]

    try:
        return _execute_direct_workspace(cached_workspace, hidden_states, routed_weights, routed_ids)
    except Exception:
        _DIRECT_BS512_EXPERT33_REGISTRY[cache_key] = None
        return None


def custom_kernel(data: input_t) -> output_t:
    (
        hidden_states,
        _raw_gate_up,
        _raw_down,
        _raw_gate_scale,
        _raw_down_scale,
        gate_up_shuffled,
        down_shuffled,
        gate_up_scale_shuffled,
        down_scale_shuffled,
        topk_weights,
        topk_ids,
        config,
    ) = data
    hidden_pad = int(config["d_hidden_pad"] - config["d_hidden"])
    inter_pad = int(config["d_expert_pad"] - config["d_expert"])

    token = int(hidden_states.shape[0])
    inter_dim = int(config["d_expert"])
    expert = int(config["n_routed_experts"] + config["n_shared_experts"])
    if _is_bs512_expert33_target(config, token=token, expert=expert, inter_dim=inter_dim):
        direct_out = _try_direct_bs512_expert33(data)
        if direct_out is not None:
            return direct_out

    return fused_moe(
        hidden_states,
        gate_up_shuffled,
        down_shuffled,
        topk_weights,
        topk_ids,
        expert_mask=None,
        activation=ActivationType.Silu,
        quant_type=QuantType.per_1x32,
        doweight_stage1=False,
        w1_scale=gate_up_scale_shuffled,
        w2_scale=down_scale_shuffled,
        a1_scale=None,
        a2_scale=None,
        hidden_pad=hidden_pad,
        intermediate_pad=inter_pad,
    )

scrolls · 704 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 683939.

⋯ 2 unchanged lines
from __future__ import annotations
import functools
+ import inspect
import math
import os
from dataclasses import dataclass
⋯ 7 unchanged lines
import aiter
import aiter.fused_moe as fmoe
import aiter.ops.flydsl.moe_kernels as flydsl_kernels
- from aiter import ActivationType, QuantType
+ from aiter import ActivationType, QuantType, dtypes
from aiter.fused_moe import fused_moe
+ from aiter.ops.moe_sorting import moe_sorting_fwd
+ from aiter.ops.triton.quant.fused_mxfp4_quant import fused_dynamic_mxfp4_quant_moe_sort
torch.set_grad_enabled(False)
⋯ 12 unchanged lines
def _register_missing_flydsl_kernels() -> None:
- # Some builds omit this key, but config rows may reference it.
flydsl_kernels._KERNEL_PARAMS.setdefault(
"flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic",
{
⋯ 42 unchanged lines
"Quant3_MulRoutedWeight0_silu_FP4X2_FP4X2_B16"
)
_K2_F16_128_A = "flydsl_moe2_afp4_wfp4_bf16_t16x128x128_atomic"
+ _K2_FLY_64_REDUCE = "flydsl_moe2_afp4_wfp4_bf16_t64x256x256_reduce"
@dataclass(frozen=True)
⋯ 46 unchanged lines
return row
+ def _retarget_stage_kernel(stage: Any, *, kernel_name: str, block_m: int) -> Any:
+ if not isinstance(stage, functools.partial):
+ return stage
+ try:
+ params = set(inspect.signature(stage.func).parameters)
+ except Exception:
+ params = set()
+ kw = dict(stage.keywords or {})
+ if "kernelName" in kw or "kernelName" in params or not params:
+ kw["kernelName"] = kernel_name
+ if "block_m" in kw or "block_m" in params or not params:
+ kw["block_m"] = block_m
+ return functools.partial(stage.func, *stage.args, **kw)
+
+
+ def _supports_kernel_retarget(stage: Any) -> bool:
+ if not isinstance(stage, functools.partial):
+ return False
+ try:
+ params = set(inspect.signature(stage.func).parameters)
+ except Exception:
+ params = set()
+ kw = dict(stage.keywords or {})
+ return ("kernelName" in params) or ("kernelName" in kw) or (not params)
+
+
def _manual_plans() -> Iterable[ShapePlan]:
model_dim = 7168
yield ShapePlan(16, model_dim, 256, 257, block_m=16, ksplit=2)
⋯ 47 unchanged lines
if col in merged.columns:
merged[col] = merged[col].fillna(0).astype(int)
- # Keep cktile for tiny utilization.
util = (merged["token"] * merged["topk"]) // merged["expert"].clip(lower=1)
merged.loc[(merged["expert"] > 0) & (merged["topk"] > 0) & (util <= 10), "ksplit"] = 4
- # Remove known weak rows; force manual rows below.
kill = (merged["expert"] == 257) & (merged["inter_dim"] == 256) & (merged["token"].isin([16, 128]))
merged = merged[~kill]
⋯ 40 unchanged lines
def _patch_runtime_nt_from_cfg() -> None:
base_get_cfg = fmoe.get_2stage_cfgs
+ enable_expert257_bs128_tuned = True
@functools.lru_cache(maxsize=2048)
def wrapped(
⋯ 51 unchanged lines
doweight_stage1,
)
row = fmoe.cfg_2stages.get(sig) if isinstance(fmoe.cfg_2stages, dict) else None
- if not row or row.get("use_non_temporal_load") is None:
- return meta
- flag = bool(row["use_non_temporal_load"])
+ out = meta
+ if row and row.get("use_non_temporal_load") is not None:
+ flag = bool(row["use_non_temporal_load"])
- def maybe_toggle(stage):
- if not isinstance(stage, functools.partial):
+ def maybe_toggle(stage):
+ if not isinstance(stage, functools.partial):
+ return stage
+ kw = dict(stage.keywords or {})
+ if "use_non_temporal_load" in kw:
+ kw["use_non_temporal_load"] = flag
+ return functools.partial(stage.func, *stage.args, **kw)
+ if "non_temporal_load" in kw:
+ kw["non_temporal_load"] = flag
+ return functools.partial(stage.func, *stage.args, **kw)
return stage
- kw = dict(stage.keywords or {})
- if "use_non_temporal_load" in kw:
- kw["use_non_temporal_load"] = flag
- return functools.partial(stage.func, *stage.args, **kw)
- if "non_temporal_load" in kw:
- kw["non_temporal_load"] = flag
- return functools.partial(stage.func, *stage.args, **kw)
- return stage
- s1 = maybe_toggle(meta.stage1)
- s2 = maybe_toggle(meta.stage2)
- if s1 is meta.stage1 and s2 is meta.stage2:
- return meta
- return fmoe.MOEMetadata(s1, s2, meta.block_m, meta.ksplit, meta.run_1stage, meta.has_bias, flag)
+ s1 = maybe_toggle(out.stage1)
+ s2 = maybe_toggle(out.stage2)
+ out = fmoe.MOEMetadata(s1, s2, out.block_m, out.ksplit, out.run_1stage, out.has_bias, flag)
+ if (
+ enable_expert257_bs128_tuned
+ and token == 128
+ and model_dim == 7168
+ and inter_dim == 256
+ and expert == 257
+ and topk == 9
+ and _supports_kernel_retarget(out.stage1)
+ and _supports_kernel_retarget(out.stage2)
+ ):
+ s1 = _retarget_stage_kernel(out.stage1, kernel_name=_K1_256x128, block_m=32)
+ s2 = _retarget_stage_kernel(out.stage2, kernel_name=_K2_FLY_64_REDUCE, block_m=32)
+ out = fmoe.MOEMetadata(s1, s2, 32, 1, False, out.has_bias, False)
+
+ return out
+
fmoe.get_2stage_cfgs = wrapped # type: ignore[assignment]
⋯ 7 unchanged lines
if hasattr(fmoe.get_2stage_cfgs, "cache_clear"):
fmoe.get_2stage_cfgs.cache_clear()
_patch_runtime_nt_from_cfg()
- # keep cheap deterministic helper cache only (not an input-result cache)
original_get_padded_m = getattr(fmoe, "get_padded_M", None)
if original_get_padded_m is not None:
fmoe.get_padded_M = functools.lru_cache(maxsize=256)(original_get_padded_m) # type: ignore[assignment]
+ def _install_expert33_dp2_sorting() -> None:
+ orig_impl = getattr(fmoe, "_moe_sorting_impl", None)
+ sort_fn = getattr(aiter, "moe_sorting_opus_fwd", None)
+ if orig_impl is None or sort_fn is None:
+ return
+
+ workspace_cache: Dict[Tuple[Any, ...], Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]] = {}
+
+ def _wrapped(
+ topk_ids,
+ topk_weights,
+ num_experts,
+ model_dim,
+ moebuf_dtype,
+ block_size,
+ expert_mask=None,
+ num_local_tokens=None,
+ dispatch_policy=0,
+ use_opus=False,
+ ):
+ if int(num_experts) != 33:
+ return orig_impl(
+ topk_ids,
+ topk_weights,
+ num_experts,
+ model_dim,
+ moebuf_dtype,
+ block_size,
+ expert_mask,
+ num_local_tokens,
+ dispatch_policy,
+ use_opus,
+ )
+
+ m_tokens, topk = topk_ids.shape
+ max_padded = int(topk_ids.numel() + num_experts * block_size - topk)
+ max_blocks = int((max_padded + block_size - 1) // block_size)
+ device = topk_ids.device
+ key = (
+ device.index if device.type == "cuda" else -1,
+ max_padded,
+ max_blocks,
+ int(m_tokens),
+ int(model_dim),
+ str(moebuf_dtype),
+ int(block_size),
+ int(num_experts),
+ )
+ cached = workspace_cache.get(key)
+ if cached is None:
+ cached = (
+ torch.empty(max_padded, dtype=torch.int32, device=device),
+ torch.empty(max_padded, dtype=torch.float32, device=device),
+ torch.empty(max_blocks, dtype=torch.int32, device=device),
+ torch.empty(2, dtype=torch.int32, device=device),
+ torch.empty((m_tokens, model_dim), dtype=moebuf_dtype, device=device),
+ )
+ workspace_cache[key] = cached
+
+ sid, sw, sei, nvi, mb = cached
+ sort_fn(
+ topk_ids,
+ topk_weights,
+ sid,
+ sw,
+ sei,
+ nvi,
+ mb,
+ num_experts,
+ block_size,
+ expert_mask,
+ num_local_tokens,
+ 2,
+ )
+ return sid, sw, sei, nvi, mb
+
+ fmoe._moe_sorting_impl = _wrapped
+
+
_bootstrap()
+ _install_expert33_dp2_sorting()
+ @dataclass
+ class _DirectBS512Expert33Workspace:
+ meta: Any
+ block_m: int
+ token: int
+ topk: int
+ inter_dim: int
+ total_experts: int
+ gate_up_shuffled: torch.Tensor
+ down_shuffled: torch.Tensor
+ w1_scale: torch.Tensor
+ w2_scale: torch.Tensor
+ sorted_ids: torch.Tensor
+ sorted_weights: torch.Tensor
+ sorted_expert_ids: torch.Tensor
+ num_valid_ids: torch.Tensor
+ moe_out: torch.Tensor
+ a2_buf: torch.Tensor
+
+
+ _DIRECT_BS512_EXPERT33_REGISTRY: Dict[Tuple[Any, ...], _DirectBS512Expert33Workspace | None] = {}
+
+
+ def _is_bs512_expert33_target(cfg: Dict[str, Any], token: int, expert: int, inter_dim: int) -> bool:
+ return (
+ token == 512
+ and expert == 33
+ and inter_dim in (512, 2048)
+ and int(cfg.get("n_routed_experts", 0)) == 32
+ and int(cfg.get("n_shared_experts", 0)) == 1
+ and int(cfg.get("d_hidden", 0)) == 7168
+ )
+
+
+ def _direct_registry_key(data: input_t) -> Tuple[Any, ...]:
+ hs, w1, w2, s1, s2, cfg = data[0], data[5], data[6], data[7], data[8], data[11]
+ dev = hs.device
+ return (
+ dev.type,
+ dev.index if dev.type == "cuda" else -1,
+ int(cfg["d_expert"]),
+ int(w1.data_ptr()),
+ int(w2.data_ptr()),
+ int(s1.data_ptr()),
+ int(s2.data_ptr()),
+ )
+
+
+ def _as_fp8_e8m0_scale(weight: torch.Tensor, scale: torch.Tensor) -> torch.Tensor:
+ return scale.view(dtypes.fp8_e8m0) if weight.dtype == dtypes.fp4x2 else scale
+
+
+ def _make_direct_workspace(data: input_t) -> _DirectBS512Expert33Workspace | None:
+ hs, w1, w2, s1, s2, tw, cfg = data[0], data[5], data[6], data[7], data[8], data[9], data[11]
+ token = int(hs.shape[0])
+ total_experts = int(w1.shape[0])
+ topk = int(tw.shape[1])
+ hidden_pad = int(cfg["d_hidden_pad"] - cfg["d_hidden"])
+ inter_pad = int(cfg["d_expert_pad"] - cfg["d_expert"])
+ _, model_dim, inter_dim = fmoe.get_inter_dim(w1.shape, w2.shape)
+
+ meta = fmoe.get_2stage_cfgs(
+ fmoe.get_padded_M(token),
+ model_dim,
+ inter_dim,
+ total_experts,
+ topk,
+ hs.dtype,
+ dtypes.fp4x2,
+ w1.dtype,
+ QuantType.per_1x32,
+ inter_dim != w1.shape[1],
+ ActivationType.Silu,
+ False,
+ hidden_pad,
+ inter_pad,
+ getattr(w1, "is_shuffled", False),
+ )
+ if meta.run_1stage:
+ return None
+
+ block_m = int(meta.block_m)
+ padded_tokens = int(token * topk + total_experts * block_m - topk)
+ padded_blocks = int((padded_tokens + block_m - 1) // block_m)
+ device = hs.device
+ return _DirectBS512Expert33Workspace(
+ meta=meta,
+ block_m=block_m,
+ token=token,
+ topk=topk,
+ inter_dim=int(inter_dim),
+ total_experts=total_experts,
+ gate_up_shuffled=w1,
+ down_shuffled=w2,
+ w1_scale=_as_fp8_e8m0_scale(w1, s1),
+ w2_scale=_as_fp8_e8m0_scale(w2, s2),
+ sorted_ids=torch.empty(padded_tokens, dtype=torch.int32, device=device),
+ sorted_weights=torch.empty(padded_tokens, dtype=torch.float32, device=device),
+ sorted_expert_ids=torch.empty(padded_blocks, dtype=torch.int32, device=device),
+ num_valid_ids=torch.empty(2, dtype=torch.int32, device=device),
+ moe_out=torch.empty((token, int(model_dim)), dtype=hs.dtype, device=device),
+ a2_buf=torch.empty((token, topk, int(inter_dim)), dtype=hs.dtype, device=device),
+ )
+
+
+ def _execute_direct_workspace(
+ ws: _DirectBS512Expert33Workspace,
+ hidden_states: torch.Tensor,
+ topk_weights: torch.Tensor,
+ topk_ids: torch.Tensor,
+ ) -> output_t:
+ moe_sorting_fwd(
+ topk_ids,
+ topk_weights,
+ ws.sorted_ids,
+ ws.sorted_weights,
+ ws.sorted_expert_ids,
+ ws.num_valid_ids,
+ ws.moe_out,
+ ws.total_experts,
+ ws.block_m,
+ None,
+ None,
+ 0,
+ )
+ q_hidden, q_hidden_scale = fused_dynamic_mxfp4_quant_moe_sort(
+ hidden_states,
+ sorted_ids=ws.sorted_ids,
+ num_valid_ids=ws.num_valid_ids,
+ token_num=ws.token,
+ topk=1,
+ block_size=ws.block_m,
+ )
+ ws.meta.stage1(
+ q_hidden,
+ ws.gate_up_shuffled,
+ ws.down_shuffled,
+ ws.sorted_ids,
+ ws.sorted_expert_ids,
+ ws.num_valid_ids,
+ ws.a2_buf,
+ ws.topk,
+ block_m=ws.block_m,
+ a1_scale=q_hidden_scale,
+ w1_scale=ws.w1_scale,
+ sorted_weights=None,
+ )
+ q_a2, q_a2_scale = fused_dynamic_mxfp4_quant_moe_sort(
+ ws.a2_buf.view(-1, ws.inter_dim),
+ sorted_ids=ws.sorted_ids,
+ num_valid_ids=ws.num_valid_ids,
+ token_num=ws.token,
+ topk=ws.topk,
+ block_size=ws.block_m,
+ )
+ ws.meta.stage2(
+ q_a2.view(ws.token, ws.topk, -1),
+ ws.gate_up_shuffled,
+ ws.down_shuffled,
+ ws.sorted_ids,
+ ws.sorted_expert_ids,
+ ws.num_valid_ids,
+ ws.moe_out,
+ ws.topk,
+ w2_scale=ws.w2_scale,
+ a2_scale=q_a2_scale,
+ block_m=ws.block_m,
+ sorted_weights=ws.sorted_weights,
+ )
+ return ws.moe_out
+
+
+ def _try_direct_bs512_expert33(data: input_t) -> output_t | None:
+ cache_key = _direct_registry_key(data)
+ try:
+ cached_workspace = _DIRECT_BS512_EXPERT33_REGISTRY[cache_key]
+ except KeyError:
+ cached_workspace = _make_direct_workspace(data)
+ _DIRECT_BS512_EXPERT33_REGISTRY[cache_key] = cached_workspace
+
+ if cached_workspace is None:
+ return None
+
+ hidden_states = data[0]
+ routed_weights = data[9]
+ routed_ids = data[10]
+
+ try:
+ return _execute_direct_workspace(cached_workspace, hidden_states, routed_weights, routed_ids)
+ except Exception:
+ _DIRECT_BS512_EXPERT33_REGISTRY[cache_key] = None
+ return None
+
+
def custom_kernel(data: input_t) -> output_t:
(
hidden_states,
⋯ 11 unchanged lines
) = data
hidden_pad = int(config["d_hidden_pad"] - config["d_hidden"])
inter_pad = int(config["d_expert_pad"] - config["d_expert"])
+
+ token = int(hidden_states.shape[0])
+ inter_dim = int(config["d_expert"])
+ expert = int(config["n_routed_experts"] + config["n_shared_experts"])
+ if _is_bs512_expert33_target(config, token=token, expert=expert, inter_dim=inter_dim):
+ direct_out = _try_direct_bs512_expert33(data)
+ if direct_out is not None:
+ return direct_out
+
return fused_moe(
hidden_states,
gate_up_shuffled,
scrolls · 456 diff lines total

Best evidence level for this revision: reported

JSON