Skip to content
KernelIndex
Search⌘K

submission 697388

.jonnss · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

Submission_v290.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-697388?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
138.4µs
#110 of 782
2026-04-02

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:1d004d6dc15c1571b164adceaa235d5797276236261e5c7029c8201390e112b6
license declaredunknown
license concludedunknown
authors.jonnss
imported2026-08-15

Techniques

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

fp4a_dtype="fp4",
split-kdef _split_kw_name(fn) -> str:

Kernel source

Submission_v290.py325 lines
import dataclasses
import functools
import gc
import inspect
import importlib
import os

import torch

from task import input_t, output_t

os.environ["VLLM_MOE_CHUNK_SIZE"] = "512"
os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
os.environ["AITER_BYPASS_TUNE_CONFIG"] = "1"
os.environ["AITER_KSPLIT"] = "2"

from aiter import ActivationType, QuantType
from aiter.fused_moe import fused_moe


gc.disable()


_PATCH_DONE = False
_CURRENT_SORT_BLOCK_SIZE = 32
_CURRENT_SORT_EXPERT = -1
_ADAPTIVE_QUANT_THRESHOLD = 1_000_000
_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,
}


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

    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 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,
            )
        except Exception:
            pass

    return meta


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"


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)


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)


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,
        mode="atomic",
        MPerBlock=16,
    )
    kernel_params[_FLYDSL_STAGE2_KEY] = base


def _precompile_flydsl_defs(moe_kernels) -> None:
    get_compiled_stage2 = getattr(moe_kernels, "_get_compiled_stage2", None)
    if get_compiled_stage2 is None:
        return

    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


def _choose_block_size_m(hidden_states: torch.Tensor, gate_up_weight_shuffled: torch.Tensor) -> int | None:
    m = int(hidden_states.shape[0])
    experts = int(gate_up_weight_shuffled.shape[0])
    if m <= 16:
        return 16
    if m <= 128 and experts > 64:
        return 32
    return None


def _desired_ksplit(token: int, expert: int) -> int:
    if token <= 16:
        return 2
    if token <= 128 and expert > 64:
        return 2
    return 0


def _install_patch() -> None:
    global _PATCH_DONE
    if _PATCH_DONE:
        return
    _PATCH_DONE = True

    try:
        aiter_mod = importlib.import_module("aiter")
        fused_moe_module = importlib.import_module("aiter.fused_moe")
        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)
        original_cfgs_sig = inspect.signature(original_cfgs) if original_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
        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)
        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):
            meta = original_cfgs(*args, **kwargs)
            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)

            ksplit = _desired_ksplit(token, expert)
            block_m = getattr(meta, "block_m", None)
            stage1 = _retune_partial(getattr(meta, "stage1", None), splitk=ksplit)
            stage2 = getattr(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)

            return _replace_meta(meta, block_m=block_m, ksplit=ksplit, stage1=stage1, stage2=stage2)

        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),
        ):
            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
    except Exception:
        pass


@torch.inference_mode()
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)

    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,
    )
scrolls · 325 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 692250.

+ import dataclasses
import functools
+ import gc
+ import inspect
import importlib
import os
⋯ 2 unchanged lines
from task import input_t, output_t
os.environ["VLLM_MOE_CHUNK_SIZE"] = "512"
+ os.environ["HIP_FORCE_DEV_KERNARG"] = "1"
+ os.environ["AITER_BYPASS_TUNE_CONFIG"] = "1"
+ os.environ["AITER_KSPLIT"] = "2"
from aiter import ActivationType, QuantType
from aiter.fused_moe import fused_moe
+ gc.disable()
+
+
_PATCH_DONE = False
+ _CURRENT_SORT_BLOCK_SIZE = 32
+ _CURRENT_SORT_EXPERT = -1
+ _ADAPTIVE_QUANT_THRESHOLD = 1_000_000
+ _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,
+ }
+ 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
+
+ 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 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,
+ )
+ except Exception:
+ pass
+
+ return meta
+
+
+ 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"
+
+
+ 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)
+
+
+ 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)
+
+
+ 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,
+ mode="atomic",
+ MPerBlock=16,
+ )
+ kernel_params[_FLYDSL_STAGE2_KEY] = base
+
+
+ def _precompile_flydsl_defs(moe_kernels) -> None:
+ get_compiled_stage2 = getattr(moe_kernels, "_get_compiled_stage2", None)
+ if get_compiled_stage2 is None:
+ return
+
+ 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
+
+
+ def _choose_block_size_m(hidden_states: torch.Tensor, gate_up_weight_shuffled: torch.Tensor) -> int | None:
+ m = int(hidden_states.shape[0])
+ experts = int(gate_up_weight_shuffled.shape[0])
+ if m <= 16:
+ return 16
+ if m <= 128 and experts > 64:
+ return 32
+ return None
+
+
+ def _desired_ksplit(token: int, expert: int) -> int:
+ if token <= 16:
+ return 2
+ if token <= 128 and expert > 64:
+ return 2
+ return 0
+
+
def _install_patch() -> None:
global _PATCH_DONE
if _PATCH_DONE:
⋯ 1 unchanged lines
_PATCH_DONE = True
try:
+ aiter_mod = importlib.import_module("aiter")
fused_moe_module = importlib.import_module("aiter.fused_moe")
- original = getattr(fused_moe_module, "get_2stage_cfgs", None)
- if original is None:
- return
+ fp4_utils = importlib.import_module("aiter.utility.fp4_utils")
+ moe_kernels = importlib.import_module("aiter.ops.flydsl.moe_kernels")
- def patched_get_2stage_cfgs(*args, **kwargs):
- meta = original(*args, **kwargs)
+ _register_flydsl_t16(moe_kernels)
+ _precompile_flydsl_defs(moe_kernels)
- token = kwargs.get("token", args[0] if len(args) > 0 else None)
- model_dim = kwargs.get("model_dim", args[1] if len(args) > 1 else None)
- inter_dim = kwargs.get("inter_dim", args[2] if len(args) > 2 else None)
- expert = kwargs.get("expert", args[3] if len(args) > 3 else None)
- topk = kwargs.get("topk", args[4] if len(args) > 4 else None)
+ original_cfgs = getattr(fused_moe_module, "get_2stage_cfgs", None)
+ original_cfgs_sig = inspect.signature(original_cfgs) if original_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
+ 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)
+ if original_cfgs is None or flydsl_stage2_wrapper is None:
+ return
- if token == 512 and model_dim == 7168 and inter_dim == 512 and expert == 33 and topk == 9:
+ def _extract_int(
+ sig: inspect.Signature | None,
+ call_args,
+ call_kwargs,
+ name: str,
+ fallback_index: int,
+ ) -> int:
+ if sig is not None:
try:
- meta.block_m = 128
+ bound_args = sig.bind_partial(*call_args, **call_kwargs).arguments
+ if name in bound_args:
+ return int(bound_args[name])
except Exception:
pass
- for attr in ("stage1", "stage2"):
- fn = getattr(meta, attr, None)
- if isinstance(fn, functools.partial):
- keywords = dict(fn.keywords or {})
- keywords["block_m"] = 128
- setattr(meta, attr, functools.partial(fn.func, *(fn.args or ()), **keywords))
- return meta
+ 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):
+ meta = original_cfgs(*args, **kwargs)
+ 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)
+
+ ksplit = _desired_ksplit(token, expert)
+ block_m = getattr(meta, "block_m", None)
+ stage1 = _retune_partial(getattr(meta, "stage1", None), splitk=ksplit)
+ stage2 = getattr(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)
+
+ return _replace_meta(meta, block_m=block_m, ksplit=ksplit, stage1=stage1, stage2=stage2)
+
+ 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
- fused_moe_globals = getattr(fused_moe, "__globals__", None)
- if isinstance(fused_moe_globals, dict) and fused_moe_globals.get("get_2stage_cfgs") is original:
- fused_moe_globals["get_2stage_cfgs"] = patched_get_2stage_cfgs
+ for maybe_globals in (
+ getattr(fused_moe, "__globals__", None),
+ getattr(fused_moe_internal, "__globals__", None),
+ getattr(fused_moe_2stages, "__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
except Exception:
pass
⋯ 19 unchanged lines
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)
return fused_moe(
hidden_states,
⋯ 9 unchanged lines
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,
)
scrolls · 319 diff lines total

Best evidence level for this revision: reported

JSON