Skip to content
KernelIndex
Search⌘K

submission 754731

.jonnss · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:7816cbaf712d98a692ff897c527a5d6471d231f7ce2a96c6ff7fd37dc589b794
license declaredunknown
license concludedunknown
authors.jonnss
imported2026-08-15

Techniques

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

fp4MoE MXFP4 two-stage kernel — MI355X

Kernel source

v84_v2_sub.py297 lines
"""
MoE MXFP4 two-stage kernel — MI355X
Patches aiter in-place so FlyDSL handles both GEMM stages for E=33 shapes.
E=257 shapes continue to use the CK fallback path.
"""

import gc
import os
import re
import shutil
import subprocess
import sys

# --- runtime knobs ---
os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")
os.environ.setdefault("HSA_ENABLE_INTERRUPT", "0")

# ------------------------------------------------------------------ #
#  Step 1 – bring in a recent flydsl wheel so the MLIR compiler works #
# ------------------------------------------------------------------ #
_WHEEL_DIR = "/tmp/_flydsl_wheel"
os.makedirs(_WHEEL_DIR, exist_ok=True)
subprocess.run(
    [sys.executable, "-m", "pip", "install", "flydsl==0.1.1",
     "--target", _WHEEL_DIR, "--upgrade", "--break-system-packages"],
    capture_output=True, timeout=120,
)
sys.path.insert(0, _WHEEL_DIR)
# flush any stale cached module objects
for _k in [k for k in sys.modules if "flydsl" in k]:
    del sys.modules[_k]

# ------------------------------------------------------------------ #
#  Step 2 – pull a recent aiter checkout and copy the FlyDSL files   #
# ------------------------------------------------------------------ #
_AITER_ROOT  = "/home/runner/aiter"
_CLONE_DIR   = "/tmp/_aiter_clone"

# reset any local edits from previous runs
subprocess.run(["git", "checkout", "--", "."],
               capture_output=True, cwd=_AITER_ROOT, timeout=15)

if not os.path.isdir(os.path.join(_CLONE_DIR, "aiter")):
    subprocess.run(
        ["git", "clone", "--depth=50",
         "https://github.com/ROCm/aiter.git", _CLONE_DIR],
        capture_output=True, timeout=120,
    )

# make sure the clone has _get_compiled_stage1 – pin to a commit that does
_kernels_py = os.path.join(_CLONE_DIR, "aiter", "ops", "flydsl", "moe_kernels.py")
if os.path.exists(_kernels_py):
    with open(_kernels_py) as _fh:
        if "_get_compiled_stage1" not in _fh.read():
            _log = subprocess.run(
                ["git", "log", "--format=%H", "-50"],
                capture_output=True, text=True, cwd=_CLONE_DIR, timeout=15,
            )
            for _sha in _log.stdout.strip().split("\n")[1:]:
                _sha = _sha.strip()
                if not _sha:
                    continue
                subprocess.run(["git", "checkout", _sha],
                               capture_output=True, cwd=_CLONE_DIR, timeout=15)
                with open(_kernels_py) as _fh2:
                    if "_get_compiled_stage1" in _fh2.read():
                        break

# copy the updated fused_moe + FlyDSL kernel files into the live aiter tree
shutil.copy2(
    os.path.join(_CLONE_DIR, "aiter", "fused_moe.py"),
    os.path.join(_AITER_ROOT, "aiter", "fused_moe.py"),
)

_src_flydsl = os.path.join(_CLONE_DIR, "aiter", "ops", "flydsl")
_dst_flydsl = os.path.join(_AITER_ROOT, "aiter", "ops", "flydsl")
for _fname in ("moe_kernels.py", "utils.py"):
    shutil.copy2(os.path.join(_src_flydsl, _fname),
                 os.path.join(_dst_flydsl, _fname))

# patch out the version guard in __init__.py
_init = os.path.join(_dst_flydsl, "__init__.py")
with open(_init) as _fh:
    _txt = _fh.read()
_txt = re.sub(r"raise ImportError\([^)]*\)", "pass  # version check skipped", _txt)
with open(_init, "w") as _fh:
    _fh.write(_txt)

# copy kernel MLIR sources if present
_src_k = os.path.join(_src_flydsl, "kernels")
_dst_k = os.path.join(_dst_flydsl, "kernels")
if os.path.isdir(_src_k):
    os.makedirs(_dst_k, exist_ok=True)
    for _fname in os.listdir(_src_k):
        if _fname.endswith(".py"):
            shutil.copy2(os.path.join(_src_k, _fname),
                         os.path.join(_dst_k, _fname))

# ------------------------------------------------------------------ #
#  Step 3 – source-patch fused_moe.py                                #
#  Inject FlyDSL dispatch functions and override get_2stage_cfgs.    #
# ------------------------------------------------------------------ #
_fmoe_path = os.path.join(_AITER_ROOT, "aiter", "fused_moe.py")
with open(_fmoe_path) as _fh:
    _fmoe_src = _fh.read()

_patch = r'''

# ---- FlyDSL stage-1: gate+up+SiLU (E<=100 shapes only) ----
def _moe_stage1_fly(hidden_states, w1, w2, sorted_token_ids,
                    sorted_expert_ids, num_valid_ids, out, topk,
                    w1_scale=None, a1_scale=None, sorted_weights=None, **_kw):
    from aiter.ops.flydsl.moe_kernels import flydsl_moe_stage1
    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=32, tile_n=256, tile_k=128,
        a_dtype="fp4", b_dtype="fp4", out_dtype="bf16",
        act="silu",
        w1_scale=w1_scale, a1_scale=a1_scale,
        sorted_weights=sorted_weights,
    )
    return out

# ---- FlyDSL stage-2: down-projection ----
def _moe_stage2_fly(inter_states, w1, w2, sorted_token_ids,
                    sorted_expert_ids, num_valid_ids, out, topk,
                    w2_scale=None, a2_scale=None, sorted_weights=None, **_kw):
    from aiter.ops.flydsl.moe_kernels import flydsl_moe_stage2
    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=32, tile_n=256, tile_k=128,
        a_dtype="fp4", b_dtype="fp4", out_dtype="bf16",
        mode="atomic",
        w2_scale=w2_scale, a2_scale=a2_scale,
        sorted_weights=sorted_weights,
    )

_orig_g2sc = get_2stage_cfgs
try:
    _orig_g2sc.cache_clear()
except AttributeError:
    pass

_stage1_ready = [False]
_stage2_ready = [False]

import functools as _ft

@_ft.lru_cache(maxsize=2048)
def _patched_g2sc(*args, **kwargs):
    meta = _orig_g2sc(*args, **kwargs)
    n_experts = args[4] if len(args) >= 5 else 0
    if (_stage1_ready[0] and n_experts <= 100
            and not meta.run_1stage and meta.stage1 is not None):
        meta.stage1 = _ft.partial(_moe_stage1_fly)
    if (_stage2_ready[0] and n_experts <= 100 and not meta.run_1stage
            and meta.stage2 is not None):
        meta.stage2 = _ft.partial(_moe_stage2_fly)
    return meta

get_2stage_cfgs = _patched_g2sc
'''

_fmoe_src += _patch
with open(_fmoe_path, "w") as _fh:
    _fh.write(_fmoe_src)

# ------------------------------------------------------------------ #
#  Step 4 – fix Triton constexpr operator signatures if needed       #
# ------------------------------------------------------------------ #
try:
    import triton.language.core as _tlc
    _ce = _tlc.constexpr
    try:
        _ce(1).__lt__(_ce(2), _semantic=None)
    except TypeError:
        _ops = [
            "__lt__", "__le__", "__gt__", "__ge__", "__eq__", "__ne__",
            "__add__", "__radd__", "__sub__", "__rsub__", "__mul__", "__rmul__",
            "__truediv__", "__floordiv__", "__mod__", "__pow__",
            "__lshift__", "__rshift__", "__and__", "__or__", "__xor__",
        ]
        for _op in _ops:
            _orig_op = getattr(_ce, _op, None)
            if _orig_op and not getattr(_orig_op, "_patched", False):
                def _wrap(fn):
                    def _inner(self, *a, _semantic=None, **kw):
                        return fn(self, *a, **kw)
                    _inner._patched = True
                    return _inner
                setattr(_ce, _op, _wrap(_orig_op))
    del _ce, _tlc
except Exception:
    pass

# ------------------------------------------------------------------ #
#  Step 5 – normal imports (see patched aiter on disk)               #
# ------------------------------------------------------------------ #
import torch

from task import input_t, output_t
from aiter import ActivationType, QuantType
from aiter.fused_moe import fused_moe
import aiter.fused_moe as _fm

# ------------------------------------------------------------------ #
#  Step 6 – pre-compile FlyDSL MLIR + set ready flags               #
# ------------------------------------------------------------------ #
try:
    from aiter.ops.flydsl.moe_kernels import (
        _get_compiled_stage1,
        _get_compiled_stage2,
    )

    # compile stage-1 for both E=33 inter-dim variants
    for _idim in (512, 2048):
        try:
            _get_compiled_stage1(
                model_dim=7168, inter_dim=_idim, experts=33, topk=9,
                tile_m=32, tile_n=256, tile_k=128,
                doweight=True, a_dtype="fp4", b_dtype="fp4",
                out_dtype="bf16", act="silu",
            )
        except Exception as _exc:
            print(f"[kernel] stage1 compile failed (inter={_idim}): {_exc}",
                  file=sys.stderr)

    _fm._stage1_ready[0] = True

    # compile stage-2 for all relevant shapes
    for _idim, _nexp in [(256, 257), (512, 33), (2048, 33)]:
        try:
            _get_compiled_stage2(
                model_dim=7168, inter_dim=_idim, experts=_nexp, topk=9,
                tile_m=32, tile_n=256, tile_k=128,
                doweight=True, a_dtype="fp4", b_dtype="fp4",
                out_dtype="bf16", accumulate=True,
            )
        except Exception:
            pass

    _fm._stage2_ready[0] = True
    print("[kernel] FlyDSL compilation done", file=sys.stderr)

except Exception as _exc:
    print(f"[kernel] FlyDSL setup error: {_exc}", file=sys.stderr)

# ------------------------------------------------------------------ #
#  Step 7 – per-shape warmup call                                    #
# ------------------------------------------------------------------ #
_seen_shapes: set = set()

def _warmup(data):
    (hs, _guw, _dw, _guws, _dws,
     guw_sh, dw_sh, guws_sh, dws_sh, tw, ti, cfg) = data
    shape_key = (cfg["bs"], cfg["n_routed_experts"], cfg["d_expert"])
    if shape_key in _seen_shapes:
        return
    _seen_shapes.add(shape_key)
    try:
        fused_moe(
            hs[:2], guw_sh, dw_sh, tw[:2], ti[:2],
            activation=ActivationType.Silu,
            quant_type=QuantType.per_1x32,
            w1_scale=guws_sh, w2_scale=dws_sh,
            hidden_pad=int(cfg["d_hidden_pad"]) - int(cfg["d_hidden"]),
            intermediate_pad=int(cfg["d_expert_pad"]) - int(cfg["d_expert"]),
        )
        torch.cuda.synchronize()
    except Exception:
        pass



# ------------------------------------------------------------------ #
#  Kernel entry point                                                 #
# ------------------------------------------------------------------ #
def custom_kernel(data: input_t) -> output_t:
    _warmup(data)
    (hs, _guw, _dw, _guws, _dws,
     guw_sh, dw_sh, guws_sh, dws_sh, tw, ti, cfg) = data
    return fused_moe(
        hs, guw_sh, dw_sh, tw, ti,
        activation=ActivationType.Silu,
        quant_type=QuantType.per_1x32,
        w1_scale=guws_sh,
        w2_scale=dws_sh,
        hidden_pad=int(cfg["d_hidden_pad"]) - int(cfg["d_hidden"]),
        intermediate_pad=int(cfg["d_expert_pad"]) - int(cfg["d_expert"]),
    )
scrolls · 297 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 715459.

- import dataclasses
- import functools
+ """
+ MoE MXFP4 two-stage kernel — MI355X
+ Patches aiter in-place so FlyDSL handles both GEMM stages for E=33 shapes.
+ E=257 shapes continue to use the CK fallback path.
+ """
+
import gc
- import inspect
- import importlib
import os
+ import re
+ import shutil
+ import subprocess
+ import sys
- import torch
+ # --- runtime knobs ---
+ os.environ.setdefault("HIP_FORCE_DEV_KERNARG", "1")
+ os.environ.setdefault("HSA_ENABLE_INTERRUPT", "0")
- from task import input_t, output_t
-
- os.environ["VLLM_MOE_CHUNK_SIZE"] = "512"
- os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
-
- from aiter import ActivationType, QuantType
- from aiter.fused_moe import fused_moe
-
-
- gc.disable()
-
-
- _PATCH_DONE = False
- _DIRECT_FUSED_MOE = None
- _FUSED_MOE_MODULE = None
- _CURRENT_SORT_BLOCK_SIZE = 32
- _CURRENT_SORT_EXPERT = -1
- _ADAPTIVE_QUANT_THRESHOLD = 1_000_000
- _CFG_CACHE = {}
- _SORT_BLOCK_CACHE = {}
- _FLYDSL_STAGE2_KEY = "flydsl_moe2_afp4_wfp4_bf16_t16x256x128_atomic"
- _FLYDSL_BASE_KEY = "flydsl_moe2_afp4_wfp4_bf16_t32x256x256_atomic"
- _E33_STAGE2_SHAPES = {
- (128, 7168, 512, 33, 9): 64,
- (512, 7168, 512, 33, 9): 128,
- (512, 7168, 2048, 33, 9): 64,
- }
- _DIRECT_PATH_SHAPES = frozenset(
- {
- (512, 7168, 256, 257, 9),
- (128, 7168, 512, 33, 9),
- (512, 7168, 512, 33, 9),
- (512, 7168, 2048, 33, 9),
- }
+ # ------------------------------------------------------------------ #
+ # Step 1 – bring in a recent flydsl wheel so the MLIR compiler works #
+ # ------------------------------------------------------------------ #
+ _WHEEL_DIR = "/tmp/_flydsl_wheel"
+ os.makedirs(_WHEEL_DIR, exist_ok=True)
+ subprocess.run(
+ [sys.executable, "-m", "pip", "install", "flydsl==0.1.1",
+ "--target", _WHEEL_DIR, "--upgrade", "--break-system-packages"],
+ capture_output=True, timeout=120,
)
+ sys.path.insert(0, _WHEEL_DIR)
+ # flush any stale cached module objects
+ for _k in [k for k in sys.modules if "flydsl" in k]:
+ del sys.modules[_k]
+ # ------------------------------------------------------------------ #
+ # Step 2 – pull a recent aiter checkout and copy the FlyDSL files #
+ # ------------------------------------------------------------------ #
+ _AITER_ROOT = "/home/runner/aiter"
+ _CLONE_DIR = "/tmp/_aiter_clone"
- def _replace_meta(meta, *, block_m=None, ksplit=None, stage1=None, stage2=None):
- try:
- if dataclasses.is_dataclass(meta):
- updates = {}
- if block_m is not None:
- updates["block_m"] = block_m
- if ksplit is not None:
- updates["ksplit"] = ksplit
- if stage1 is not None:
- updates["stage1"] = stage1
- if stage2 is not None:
- updates["stage2"] = stage2
- if updates:
- return dataclasses.replace(meta, **updates)
- except Exception:
- pass
+ # reset any local edits from previous runs
+ subprocess.run(["git", "checkout", "--", "."],
+ capture_output=True, cwd=_AITER_ROOT, timeout=15)
- try:
- if block_m is not None:
- meta.block_m = block_m
- if ksplit is not None:
- meta.ksplit = ksplit
- if stage1 is not None:
- meta.stage1 = stage1
- if stage2 is not None:
- meta.stage2 = stage2
- return meta
- except Exception:
- pass
+ if not os.path.isdir(os.path.join(_CLONE_DIR, "aiter")):
+ subprocess.run(
+ ["git", "clone", "--depth=50",
+ "https://github.com/ROCm/aiter.git", _CLONE_DIR],
+ capture_output=True, timeout=120,
+ )
- if all(hasattr(meta, name) for name in ("stage1", "stage2", "block_m", "ksplit", "run_1stage")):
- has_bias = getattr(meta, "has_bias", False)
- use_non_temporal_load = getattr(meta, "use_non_temporal_load", True)
- try:
- return type(meta)(
- meta.stage1 if stage1 is None else stage1,
- meta.stage2 if stage2 is None else stage2,
- meta.block_m if block_m is None else block_m,
- meta.ksplit if ksplit is None else ksplit,
- meta.run_1stage,
- has_bias,
- use_non_temporal_load,
+ # make sure the clone has _get_compiled_stage1 – pin to a commit that does
+ _kernels_py = os.path.join(_CLONE_DIR, "aiter", "ops", "flydsl", "moe_kernels.py")
+ if os.path.exists(_kernels_py):
+ with open(_kernels_py) as _fh:
+ if "_get_compiled_stage1" not in _fh.read():
+ _log = subprocess.run(
+ ["git", "log", "--format=%H", "-50"],
+ capture_output=True, text=True, cwd=_CLONE_DIR, timeout=15,
)
- except Exception:
- pass
+ for _sha in _log.stdout.strip().split("\n")[1:]:
+ _sha = _sha.strip()
+ if not _sha:
+ continue
+ subprocess.run(["git", "checkout", _sha],
+ capture_output=True, cwd=_CLONE_DIR, timeout=15)
+ with open(_kernels_py) as _fh2:
+ if "_get_compiled_stage1" in _fh2.read():
+ break
- return meta
+ # copy the updated fused_moe + FlyDSL kernel files into the live aiter tree
+ shutil.copy2(
+ os.path.join(_CLONE_DIR, "aiter", "fused_moe.py"),
+ os.path.join(_AITER_ROOT, "aiter", "fused_moe.py"),
+ )
+ _src_flydsl = os.path.join(_CLONE_DIR, "aiter", "ops", "flydsl")
+ _dst_flydsl = os.path.join(_AITER_ROOT, "aiter", "ops", "flydsl")
+ for _fname in ("moe_kernels.py", "utils.py"):
+ shutil.copy2(os.path.join(_src_flydsl, _fname),
+ os.path.join(_dst_flydsl, _fname))
- def _split_kw_name(fn) -> str:
- if isinstance(fn, functools.partial):
- keywords = dict(fn.keywords or {})
- if "splitk" in keywords:
- return "splitk"
- if "split_k" in keywords:
- return "split_k"
- target = fn.func if isinstance(fn, functools.partial) else fn
- name = getattr(target, "__name__", "")
- if name in {"cktile_moe_gemm1", "cktile_moe_stage1"}:
- return "split_k"
- return "splitk"
+ # patch out the version guard in __init__.py
+ _init = os.path.join(_dst_flydsl, "__init__.py")
+ with open(_init) as _fh:
+ _txt = _fh.read()
+ _txt = re.sub(r"raise ImportError\([^)]*\)", "pass # version check skipped", _txt)
+ with open(_init, "w") as _fh:
+ _fh.write(_txt)
+ # copy kernel MLIR sources if present
+ _src_k = os.path.join(_src_flydsl, "kernels")
+ _dst_k = os.path.join(_dst_flydsl, "kernels")
+ if os.path.isdir(_src_k):
+ os.makedirs(_dst_k, exist_ok=True)
+ for _fname in os.listdir(_src_k):
+ if _fname.endswith(".py"):
+ shutil.copy2(os.path.join(_src_k, _fname),
+ os.path.join(_dst_k, _fname))
- def _retune_partial(fn, *, block_m=None, splitk=None):
- if not isinstance(fn, functools.partial):
- return fn
- keywords = dict(fn.keywords or {})
- if block_m is not None:
- keywords["block_m"] = block_m
- if splitk is not None:
- keywords.pop("splitk", None)
- keywords.pop("split_k", None)
- keywords[_split_kw_name(fn)] = splitk
- return functools.partial(fn.func, *(fn.args or ()), **keywords)
+ # ------------------------------------------------------------------ #
+ # Step 3 – source-patch fused_moe.py #
+ # Inject FlyDSL dispatch functions and override get_2stage_cfgs. #
+ # ------------------------------------------------------------------ #
+ _fmoe_path = os.path.join(_AITER_ROOT, "aiter", "fused_moe.py")
+ with open(_fmoe_path) as _fh:
+ _fmoe_src = _fh.read()
+ _patch = r'''
- def _wrap_flydsl_stage2(stage2_wrapper, original_stage2, kernel_name: str):
- keywords = {}
- if isinstance(original_stage2, functools.partial):
- keywords.update(original_stage2.keywords or {})
- keywords["kernelName"] = kernel_name
- return functools.partial(stage2_wrapper, **keywords)
+ # ---- FlyDSL stage-1: gate+up+SiLU (E<=100 shapes only) ----
+ def _moe_stage1_fly(hidden_states, w1, w2, sorted_token_ids,
+ sorted_expert_ids, num_valid_ids, out, topk,
+ w1_scale=None, a1_scale=None, sorted_weights=None, **_kw):
+ from aiter.ops.flydsl.moe_kernels import flydsl_moe_stage1
+ 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=32, tile_n=256, tile_k=128,
+ a_dtype="fp4", b_dtype="fp4", out_dtype="bf16",
+ act="silu",
+ w1_scale=w1_scale, a1_scale=a1_scale,
+ sorted_weights=sorted_weights,
+ )
+ return out
-
- def _register_flydsl_t16(moe_kernels) -> None:
- kernel_params = getattr(moe_kernels, "_KERNEL_PARAMS", None)
- if not isinstance(kernel_params, dict):
- return
- if _FLYDSL_STAGE2_KEY in kernel_params:
- return
-
- base = dict(kernel_params.get(_FLYDSL_BASE_KEY, {}))
- if not base:
- return
-
- base.update(
- tile_m=16,
- tile_n=256,
- tile_k=128,
+ # ---- FlyDSL stage-2: down-projection ----
+ def _moe_stage2_fly(inter_states, w1, w2, sorted_token_ids,
+ sorted_expert_ids, num_valid_ids, out, topk,
+ w2_scale=None, a2_scale=None, sorted_weights=None, **_kw):
+ from aiter.ops.flydsl.moe_kernels import flydsl_moe_stage2
+ 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=32, tile_n=256, tile_k=128,
+ a_dtype="fp4", b_dtype="fp4", out_dtype="bf16",
mode="atomic",
- MPerBlock=16,
+ w2_scale=w2_scale, a2_scale=a2_scale,
+ sorted_weights=sorted_weights,
)
- kernel_params[_FLYDSL_STAGE2_KEY] = base
+ _orig_g2sc = get_2stage_cfgs
+ try:
+ _orig_g2sc.cache_clear()
+ except AttributeError:
+ pass
- def _precompile_flydsl_defs(moe_kernels) -> None:
- get_compiled_stage2 = getattr(moe_kernels, "_get_compiled_stage2", None)
- if get_compiled_stage2 is None:
- return
+ _stage1_ready = [False]
+ _stage2_ready = [False]
- for inter_dim in (512, 2048):
- try:
- get_compiled_stage2(
- model_dim=7168,
- inter_dim=inter_dim,
- experts=33,
- topk=9,
- tile_m=16,
- tile_n=256,
- tile_k=128,
- doweight=True,
- a_dtype="fp4",
- b_dtype="fp4",
- out_dtype="bf16",
- )
- except Exception:
- pass
+ import functools as _ft
+ @_ft.lru_cache(maxsize=2048)
+ def _patched_g2sc(*args, **kwargs):
+ meta = _orig_g2sc(*args, **kwargs)
+ n_experts = args[4] if len(args) >= 5 else 0
+ if (_stage1_ready[0] and n_experts <= 100
+ and not meta.run_1stage and meta.stage1 is not None):
+ meta.stage1 = _ft.partial(_moe_stage1_fly)
+ if (_stage2_ready[0] and n_experts <= 100 and not meta.run_1stage
+ and meta.stage2 is not None):
+ meta.stage2 = _ft.partial(_moe_stage2_fly)
+ return meta
- def _choose_block_size_m(hidden_states: torch.Tensor, gate_up_weight_shuffled: torch.Tensor) -> int | None:
- m = int(hidden_states.shape[0])
- if m <= 16:
- return 16
- return None
+ get_2stage_cfgs = _patched_g2sc
+ '''
+ _fmoe_src += _patch
+ with open(_fmoe_path, "w") as _fh:
+ _fh.write(_fmoe_src)
- def _should_use_direct_path(
- hidden_states: torch.Tensor,
- gate_up_weight_shuffled: torch.Tensor,
- topk_ids: torch.Tensor,
- config,
- ) -> bool:
- shape = (
- int(hidden_states.shape[0]),
- int(config["d_hidden"]),
- int(config["d_expert"]),
- int(gate_up_weight_shuffled.shape[0]),
- int(topk_ids.shape[1]),
- )
- return shape in _DIRECT_PATH_SHAPES
+ # ------------------------------------------------------------------ #
+ # Step 4 – fix Triton constexpr operator signatures if needed #
+ # ------------------------------------------------------------------ #
+ try:
+ import triton.language.core as _tlc
+ _ce = _tlc.constexpr
+ try:
+ _ce(1).__lt__(_ce(2), _semantic=None)
+ except TypeError:
+ _ops = [
+ "__lt__", "__le__", "__gt__", "__ge__", "__eq__", "__ne__",
+ "__add__", "__radd__", "__sub__", "__rsub__", "__mul__", "__rmul__",
+ "__truediv__", "__floordiv__", "__mod__", "__pow__",
+ "__lshift__", "__rshift__", "__and__", "__or__", "__xor__",
+ ]
+ for _op in _ops:
+ _orig_op = getattr(_ce, _op, None)
+ if _orig_op and not getattr(_orig_op, "_patched", False):
+ def _wrap(fn):
+ def _inner(self, *a, _semantic=None, **kw):
+ return fn(self, *a, **kw)
+ _inner._patched = True
+ return _inner
+ setattr(_ce, _op, _wrap(_orig_op))
+ del _ce, _tlc
+ except Exception:
+ pass
+ # ------------------------------------------------------------------ #
+ # Step 5 – normal imports (see patched aiter on disk) #
+ # ------------------------------------------------------------------ #
+ import torch
- def _use_bypass(token: int, expert: int) -> bool:
- return token <= 16 or (token <= 128 and expert > 64)
+ from task import input_t, output_t
+ from aiter import ActivationType, QuantType
+ from aiter.fused_moe import fused_moe
+ import aiter.fused_moe as _fm
+ # ------------------------------------------------------------------ #
+ # Step 6 – pre-compile FlyDSL MLIR + set ready flags #
+ # ------------------------------------------------------------------ #
+ try:
+ from aiter.ops.flydsl.moe_kernels import (
+ _get_compiled_stage1,
+ _get_compiled_stage2,
+ )
- def _desired_ksplit(token: int, expert: int) -> int:
- return 2 if _use_bypass(token, expert) else 0
+ # compile stage-1 for both E=33 inter-dim variants
+ for _idim in (512, 2048):
+ try:
+ _get_compiled_stage1(
+ model_dim=7168, inter_dim=_idim, experts=33, topk=9,
+ tile_m=32, tile_n=256, tile_k=128,
+ doweight=True, a_dtype="fp4", b_dtype="fp4",
+ out_dtype="bf16", act="silu",
+ )
+ except Exception as _exc:
+ print(f"[kernel] stage1 compile failed (inter={_idim}): {_exc}",
+ file=sys.stderr)
+ _fm._stage1_ready[0] = True
- def _cacheable(value):
- if isinstance(value, (str, int, float, bool, type(None))):
- return value
- if isinstance(value, tuple):
- return tuple(_cacheable(item) for item in value)
- return repr(value)
-
-
- def _make_cfg_cache_key(sig: inspect.Signature | None, use_bypass: bool, call_args, call_kwargs):
- if sig is not None:
+ # compile stage-2 for all relevant shapes
+ for _idim, _nexp in [(256, 257), (512, 33), (2048, 33)]:
try:
- bound = sig.bind_partial(*call_args, **call_kwargs)
- return (
- use_bypass,
- tuple((name, _cacheable(value)) for name, value in bound.arguments.items()),
+ _get_compiled_stage2(
+ model_dim=7168, inter_dim=_idim, experts=_nexp, topk=9,
+ tile_m=32, tile_n=256, tile_k=128,
+ doweight=True, a_dtype="fp4", b_dtype="fp4",
+ out_dtype="bf16", accumulate=True,
)
except Exception:
pass
- return (
- use_bypass,
- tuple(_cacheable(value) for value in call_args),
- tuple(sorted((key, _cacheable(value)) for key, value in call_kwargs.items())),
- )
+ _fm._stage2_ready[0] = True
+ print("[kernel] FlyDSL compilation done", file=sys.stderr)
- def _call_with_bypass_env(use_bypass: bool, fn, *args, **kwargs):
- previous = os.environ.get("AITER_BYPASS_TUNE_CONFIG")
- os.environ["AITER_BYPASS_TUNE_CONFIG"] = "1" if use_bypass else "0"
- try:
- return fn(*args, **kwargs)
- finally:
- if previous is None:
- os.environ.pop("AITER_BYPASS_TUNE_CONFIG", None)
- else:
- os.environ["AITER_BYPASS_TUNE_CONFIG"] = previous
+ except Exception as _exc:
+ print(f"[kernel] FlyDSL setup error: {_exc}", file=sys.stderr)
+ # ------------------------------------------------------------------ #
+ # Step 7 – per-shape warmup call #
+ # ------------------------------------------------------------------ #
+ _seen_shapes: set = set()
- def _install_patch() -> None:
- global _PATCH_DONE, _DIRECT_FUSED_MOE, _FUSED_MOE_MODULE
- if _PATCH_DONE:
+ def _warmup(data):
+ (hs, _guw, _dw, _guws, _dws,
+ guw_sh, dw_sh, guws_sh, dws_sh, tw, ti, cfg) = data
+ shape_key = (cfg["bs"], cfg["n_routed_experts"], cfg["d_expert"])
+ if shape_key in _seen_shapes:
return
- _PATCH_DONE = True
-
+ _seen_shapes.add(shape_key)
try:
- aiter_mod = importlib.import_module("aiter")
- fused_moe_module = importlib.import_module("aiter.fused_moe")
- _FUSED_MOE_MODULE = fused_moe_module
- fp4_utils = importlib.import_module("aiter.utility.fp4_utils")
- moe_kernels = importlib.import_module("aiter.ops.flydsl.moe_kernels")
-
- _register_flydsl_t16(moe_kernels)
- _precompile_flydsl_defs(moe_kernels)
-
- original_cfgs = getattr(fused_moe_module, "get_2stage_cfgs", None)
- raw_cfgs = getattr(original_cfgs, "__wrapped__", original_cfgs)
- original_cfgs_sig = inspect.signature(original_cfgs) if original_cfgs is not None else None
- raw_cfgs_sig = inspect.signature(raw_cfgs) if raw_cfgs is not None else None
- original_get_ksplit = getattr(fused_moe_module, "get_ksplit", None)
- original_get_ksplit_sig = inspect.signature(original_get_ksplit) if original_get_ksplit is not None else None
- original_quant_sort = getattr(fused_moe_module, "fused_dynamic_mxfp4_quant_moe_sort", None)
- quant_hip = getattr(aiter_mod, "per_1x32_f4_quant_hip", None)
- moe_mxfp4_sort = getattr(fp4_utils, "moe_mxfp4_sort", None)
- flydsl_stage2_wrapper = getattr(fused_moe_module, "_flydsl_stage2_wrapper", None)
- fused_moe_internal = getattr(fused_moe_module, "fused_moe_", None)
- fused_moe_2stages = getattr(fused_moe_module, "fused_moe_2stages", None)
- direct_fused_moe = None
- if fused_moe_internal is not None:
- try:
- direct_fused_moe = inspect.unwrap(fused_moe_internal)
- except Exception:
- direct_fused_moe = getattr(fused_moe_internal, "__wrapped__", None)
- if direct_fused_moe is None:
- direct_fused_moe = fused_moe_internal
- _DIRECT_FUSED_MOE = direct_fused_moe
- if original_cfgs is None or flydsl_stage2_wrapper is None:
- return
-
- def _extract_int(
- sig: inspect.Signature | None,
- call_args,
- call_kwargs,
- name: str,
- fallback_index: int,
- ) -> int:
- if sig is not None:
- try:
- bound_args = sig.bind_partial(*call_args, **call_kwargs).arguments
- if name in bound_args:
- return int(bound_args[name])
- except Exception:
- pass
- value = call_kwargs.get(name, call_args[fallback_index] if len(call_args) > fallback_index else -1)
- try:
- return int(value)
- except Exception:
- return -1
-
- def patched_get_2stage_cfgs(*args, **kwargs):
- global _CURRENT_SORT_BLOCK_SIZE, _CURRENT_SORT_EXPERT
- token = _extract_int(original_cfgs_sig, args, kwargs, "token", 0)
- model_dim = _extract_int(original_cfgs_sig, args, kwargs, "model_dim", 1)
- inter_dim = _extract_int(original_cfgs_sig, args, kwargs, "inter_dim", 2)
- expert = _extract_int(original_cfgs_sig, args, kwargs, "expert", 3)
- topk = _extract_int(original_cfgs_sig, args, kwargs, "topk", 4)
- shape = (token, model_dim, inter_dim, expert, topk)
- use_bypass = _use_bypass(token, expert)
- cache_key = _make_cfg_cache_key(raw_cfgs_sig, use_bypass, args, kwargs)
- meta = _CFG_CACHE.get(cache_key)
- if meta is None:
- raw_meta = _call_with_bypass_env(use_bypass, raw_cfgs, *args, **kwargs)
- ksplit = _desired_ksplit(token, expert)
- block_m = getattr(raw_meta, "block_m", None)
- stage1 = _retune_partial(getattr(raw_meta, "stage1", None), splitk=ksplit)
- stage2 = getattr(raw_meta, "stage2", None)
-
- override_block_m = _E33_STAGE2_SHAPES.get(shape)
- if override_block_m is not None:
- block_m = override_block_m
- stage1 = _retune_partial(stage1, block_m=override_block_m, splitk=0)
- stage2 = _retune_partial(stage2, block_m=override_block_m)
- stage2 = _wrap_flydsl_stage2(flydsl_stage2_wrapper, stage2, _FLYDSL_STAGE2_KEY)
- ksplit = 0
-
- if shape == (512, 7168, 256, 257, 9):
- ksplit = 0
- stage1 = _retune_partial(stage1, splitk=0)
-
- meta = _replace_meta(raw_meta, block_m=block_m, ksplit=ksplit, stage1=stage1, stage2=stage2)
- _CFG_CACHE[cache_key] = meta
-
- _CURRENT_SORT_BLOCK_SIZE = int(getattr(meta, "block_m", None) or 32)
- _CURRENT_SORT_EXPERT = expert
- return meta
-
- def patched_get_ksplit(*args, **kwargs):
- token = _extract_int(original_get_ksplit_sig, args, kwargs, "token", 0)
- expert = _extract_int(original_get_ksplit_sig, args, kwargs, "expert", 3)
- if token >= 0 and expert >= 0:
- return _desired_ksplit(token, expert)
- if original_get_ksplit is None:
- return 0
- return original_get_ksplit(*args, **kwargs)
-
- fused_moe_module.get_2stage_cfgs = patched_get_2stage_cfgs
- if original_get_ksplit is not None:
- fused_moe_module.get_ksplit = patched_get_ksplit
-
- for maybe_globals in (
- getattr(fused_moe, "__globals__", None),
- getattr(fused_moe_internal, "__globals__", None),
- getattr(fused_moe_2stages, "__globals__", None),
- getattr(direct_fused_moe, "__globals__", None),
- ):
- if isinstance(maybe_globals, dict) and maybe_globals.get("get_2stage_cfgs") is original_cfgs:
- maybe_globals["get_2stage_cfgs"] = patched_get_2stage_cfgs
- if isinstance(maybe_globals, dict) and maybe_globals.get("get_ksplit") is original_get_ksplit:
- maybe_globals["get_ksplit"] = patched_get_ksplit
-
- if original_quant_sort is None or quant_hip is None or moe_mxfp4_sort is None:
- return
-
- def split_quant_sort(
- x: torch.Tensor,
- sorted_ids: torch.Tensor,
- num_valid_ids: torch.Tensor,
- token_num: int,
- topk: int,
- block_size: int = 32,
- scaling_mode: str = "even",
- ):
- if scaling_mode != "even":
- return original_quant_sort(
- x,
- sorted_ids,
- num_valid_ids,
- token_num,
- topk,
- block_size=block_size,
- scaling_mode=scaling_mode,
- )
-
- quant_out = quant_hip(x, scale=None, shuffle=False)
- x_fp4, blockscale_e8m0 = quant_out[:2]
- requested_block_size = int(_CURRENT_SORT_BLOCK_SIZE or block_size)
- cache_key = (requested_block_size, int(_CURRENT_SORT_EXPERT))
- candidates = []
- for candidate in (
- _SORT_BLOCK_CACHE.get(cache_key),
- requested_block_size,
- int(block_size),
- max(requested_block_size, 128),
- 64,
- 128,
- 256,
- 512,
- ):
- if candidate and candidate not in candidates:
- candidates.append(candidate)
-
- last_error = None
- for candidate in candidates:
- try:
- blockscale_sorted = moe_mxfp4_sort(
- blockscale_e8m0,
- sorted_ids,
- num_valid_ids,
- token_num,
- candidate,
- )
- _SORT_BLOCK_CACHE[cache_key] = candidate
- return x_fp4, blockscale_sorted
- except AssertionError as exc:
- last_error = exc
-
- if last_error is not None:
- raise last_error
- return x_fp4, blockscale_e8m0
-
- def adaptive_quant_sort(*args, **kwargs):
- x = kwargs.get("x", args[0] if len(args) > 0 else None)
- scaling_mode = kwargs.get("scaling_mode", args[6] if len(args) > 6 else "even")
- if x is None or scaling_mode != "even":
- return original_quant_sort(*args, **kwargs)
- if int(x.numel()) > _ADAPTIVE_QUANT_THRESHOLD:
- return split_quant_sort(*args, **kwargs)
- return original_quant_sort(*args, **kwargs)
-
- fused_moe_module.fused_dynamic_mxfp4_quant_moe_sort = adaptive_quant_sort
- for maybe_globals in (
- getattr(fused_moe_internal, "__globals__", None),
- getattr(fused_moe_2stages, "__globals__", None),
- getattr(direct_fused_moe, "__globals__", None),
- ):
- if (
- isinstance(maybe_globals, dict)
- and maybe_globals.get("fused_dynamic_mxfp4_quant_moe_sort") is original_quant_sort
- ):
- maybe_globals["fused_dynamic_mxfp4_quant_moe_sort"] = adaptive_quant_sort
- except Exception:
- pass
-
-
- def _call_moe(
- hidden_states: torch.Tensor,
- gate_up_weight_shuffled: torch.Tensor,
- down_weight_shuffled: torch.Tensor,
- topk_weights: torch.Tensor,
- topk_ids: torch.Tensor,
- gate_up_weight_scale_shuffled: torch.Tensor,
- down_weight_scale_shuffled: torch.Tensor,
- hidden_pad: int,
- intermediate_pad: int,
- block_size_m: int | None,
- ):
- direct_fused_moe = _DIRECT_FUSED_MOE
- if direct_fused_moe is None:
- return fused_moe(
- hidden_states,
- gate_up_weight_shuffled,
- down_weight_shuffled,
- topk_weights,
- topk_ids,
- expert_mask=None,
+ fused_moe(
+ hs[:2], guw_sh, dw_sh, tw[:2], ti[:2],
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,
- block_size_M=block_size_m,
- hidden_pad=hidden_pad,
- intermediate_pad=intermediate_pad,
+ w1_scale=guws_sh, w2_scale=dws_sh,
+ hidden_pad=int(cfg["d_hidden_pad"]) - int(cfg["d_hidden"]),
+ intermediate_pad=int(cfg["d_expert_pad"]) - int(cfg["d_expert"]),
)
-
- block_size_arg = -1 if not block_size_m else int(block_size_m)
- try:
- return direct_fused_moe(
- hidden_states=hidden_states,
- w1=gate_up_weight_shuffled,
- w2=down_weight_shuffled,
- topk_weight=topk_weights,
- topk_ids=topk_ids,
- expert_mask=None,
- activation=ActivationType.Silu.value,
- quant_type=QuantType.per_1x32.value,
- doweight_stage1=False,
- w1_scale=gate_up_weight_scale_shuffled,
- w2_scale=down_weight_scale_shuffled,
- a1_scale=None,
- a2_scale=None,
- block_size_M=block_size_arg,
- num_local_tokens=None,
- moe_sorting_dispatch_policy=0,
- dtype=None,
- hidden_pad=hidden_pad,
- intermediate_pad=intermediate_pad,
- bias1=None,
- bias2=None,
- )
+ torch.cuda.synchronize()
except Exception:
- 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,
- block_size_M=block_size_m,
- hidden_pad=hidden_pad,
- intermediate_pad=intermediate_pad,
- )
+ pass
- def _call_moe_inlined(
- hidden_states: torch.Tensor,
- gate_up_weight_shuffled: torch.Tensor,
- down_weight_shuffled: torch.Tensor,
- topk_weights: torch.Tensor,
- topk_ids: torch.Tensor,
- gate_up_weight_scale_shuffled: torch.Tensor,
- down_weight_scale_shuffled: torch.Tensor,
- hidden_pad: int,
- intermediate_pad: int,
- block_size_m: int | None,
- ):
- fused_moe_module = _FUSED_MOE_MODULE
- if fused_moe_module is None:
- return _call_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,
- block_size_m,
- )
- try:
- get_inter_dim = getattr(fused_moe_module, "get_inter_dim")
- dtypes = getattr(fused_moe_module, "dtypes")
- quant_remap = getattr(fused_moe_module, "quant_remap")
- get_gfx = getattr(fused_moe_module, "get_gfx")
- get_padded_M = getattr(fused_moe_module, "get_padded_M")
- get_2stage_cfgs = getattr(fused_moe_module, "get_2stage_cfgs")
- moe_sorting = getattr(fused_moe_module, "moe_sorting")
- fused_moe_2stages = getattr(fused_moe_module, "fused_moe_2stages")
-
- activation = ActivationType.Silu
- quant_type = QuantType.per_1x32
- M, topk = topk_ids.shape
- E, model_dim, inter_dim = get_inter_dim(
- gate_up_weight_shuffled.shape,
- down_weight_shuffled.shape,
- )
-
- assert gate_up_weight_shuffled.shape[1] in [inter_dim, inter_dim * 2]
- is_g1u1 = inter_dim != gate_up_weight_shuffled.shape[1]
- is_shuffled = getattr(gate_up_weight_shuffled, "is_shuffled", False)
- dtype = hidden_states.dtype
- global_e = E
- q_dtype_w = gate_up_weight_shuffled.dtype
- q_dtype_a = gate_up_weight_shuffled.dtype if gate_up_weight_shuffled.dtype != torch.uint32 else dtypes.fp8
- quant_type = quant_remap.get(quant_type, quant_type)
- if quant_type == QuantType.per_1x32:
- if activation == ActivationType.Swiglu:
- if get_gfx() != "gfx950" or M < 512:
- q_dtype_a = dtypes.bf16
- else:
- q_dtype_a = dtypes.fp8
- else:
- q_dtype_a = dtypes.fp4x2
-
- metadata = get_2stage_cfgs(
- get_padded_M(M),
- model_dim,
- inter_dim,
- E,
- topk,
- dtype,
- q_dtype_a,
- q_dtype_w,
- quant_type,
- is_g1u1,
- activation,
- False,
- hidden_pad,
- intermediate_pad,
- is_shuffled,
- )
-
- block_size_eff = getattr(metadata, "block_m", None) if block_size_m is None else block_size_m
- if block_size_eff is not None:
- block_size_eff = int(block_size_eff)
-
- sorted_ids, sorted_weights, sorted_expert_ids, num_valid_ids, moe_buf = moe_sorting(
- topk_ids,
- topk_weights,
- global_e,
- model_dim,
- dtype,
- block_size_eff,
- None,
- None,
- 0,
- )
-
- if getattr(metadata, "run_1stage", False):
- return metadata.stage1(
- hidden_states,
- gate_up_weight_shuffled,
- down_weight_shuffled,
- topk,
- sorted_ids,
- sorted_weights,
- sorted_expert_ids,
- num_valid_ids,
- moe_buf,
- is_g1u1,
- block_size_eff,
- q_dtype_a=q_dtype_a,
- q_dtype_w=q_dtype_w,
- w1_scale=gate_up_weight_scale_shuffled,
- w2_scale=down_weight_scale_shuffled,
- a1_scale=None,
- a2_scale=None,
- num_local_tokens=None,
- M=M,
- device=topk_ids.device,
- doweight_stage1=False,
- )
-
- return fused_moe_2stages(
- hidden_states,
- gate_up_weight_shuffled,
- down_weight_shuffled,
- topk,
- sorted_ids,
- sorted_weights,
- sorted_expert_ids,
- num_valid_ids,
- moe_buf,
- is_g1u1,
- block_size_eff,
- activation=activation,
- quant_type=quant_type,
- doweight_stage1=False,
- q_dtype_a=q_dtype_a,
- q_dtype_w=q_dtype_w,
- w1_scale=gate_up_weight_scale_shuffled,
- w2_scale=down_weight_scale_shuffled,
- a1_scale=None,
- a2_scale=None,
- num_local_tokens=None,
- hidden_pad=hidden_pad,
- intermediate_pad=intermediate_pad,
- bias1=None,
- bias2=None,
- )
- except Exception:
- return _call_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,
- block_size_m,
- )
-
-
- @torch.inference_mode()
+ # ------------------------------------------------------------------ #
+ # Kernel entry point #
+ # ------------------------------------------------------------------ #
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
-
- _install_patch()
-
- hidden_pad = int(config["d_hidden_pad"]) - int(config["d_hidden"])
- intermediate_pad = int(config["d_expert_pad"]) - int(config["d_expert"])
- block_size_m = _choose_block_size_m(hidden_states, gate_up_weight_shuffled)
-
- if not _should_use_direct_path(hidden_states, gate_up_weight_shuffled, topk_ids, config):
- 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,
- block_size_M=block_size_m,
- hidden_pad=hidden_pad,
- intermediate_pad=intermediate_pad,
- )
-
- return _call_moe_inlined(
- 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,
- block_size_m,
+ _warmup(data)
+ (hs, _guw, _dw, _guws, _dws,
+ guw_sh, dw_sh, guws_sh, dws_sh, tw, ti, cfg) = data
+ return fused_moe(
+ hs, guw_sh, dw_sh, tw, ti,
+ activation=ActivationType.Silu,
+ quant_type=QuantType.per_1x32,
+ w1_scale=guws_sh,
+ w2_scale=dws_sh,
+ hidden_pad=int(cfg["d_hidden_pad"]) - int(cfg["d_hidden"]),
+ intermediate_pad=int(cfg["d_expert_pad"]) - int(cfg["d_expert"]),
)
scrolls · 989 diff lines total

Best evidence level for this revision: reported

JSON