Skip to content
KernelIndex
Search⌘K

submission 683939

NinoHeather · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

my_submission_refact.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-683939?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
125.2µs
#67 of 782
2026-04-01

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:806c0963b6d97e2d95b426df35250d1b62e37fede58b155f297311427b015e67
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_refact.py380 lines
#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
from __future__ import annotations

import functools
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
from aiter.fused_moe import fused_moe


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:
    # Some builds omit this key, but config rows may reference it.
    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"


@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 _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)

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

    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

    @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
        if not row or row.get("use_non_temporal_load") is None:
            return meta
        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(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)

    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()
    # 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]


_bootstrap()


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"])
    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 · 380 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 679499.

#!POPCORN leaderboard amd-moe-mxfp4
#!POPCORN gpu MI355X
-
from __future__ import annotations
+ import functools
+ 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
from aiter.fused_moe import fused_moe
+
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:
+ # Some builds omit this key, but config rows may reference it.
+ 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"
+
+
+ @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 _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)
+
+ # 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]
+
+ 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
+
+ @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
+ if not row or row.get("use_non_temporal_load") is None:
+ return meta
+ 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(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)
+
+ 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()
+ # 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]
+
+
+ _bootstrap()
+
+
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,
+ _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"])
- intermediate_pad = int(config["d_expert_pad"] - config["d_expert"])
-
+ inter_pad = int(config["d_expert_pad"] - config["d_expert"])
return fused_moe(
hidden_states,
- gate_up_weight_shuffled,
- down_weight_shuffled,
+ 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_weight_scale_shuffled,
- w2_scale=down_weight_scale_shuffled,
+ w1_scale=gate_up_scale_shuffled,
+ w2_scale=down_scale_shuffled,
a1_scale=None,
a2_scale=None,
hidden_pad=hidden_pad,
- intermediate_pad=intermediate_pad,
+ intermediate_pad=inter_pad,
)
+
scrolls · 396 diff lines total

Best evidence level for this revision: reported

JSON