Skip to content
KernelIndex
Search⌘K

submission 715459

.jonnss · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

Submission_v333.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-amd-moe-mxfp4-715459?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
127.8µs
#81 of 782
2026-04-03

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:f49a282ca38d6df3d53f57b8c3a331f3e9736934670acf8490f2d7c6f4d81b8a
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_v333.py746 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"

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),
    }
)


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])
    if m <= 16:
        return 16
    return None


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


def _use_bypass(token: int, expert: int) -> bool:
    return token <= 16 or (token <= 128 and expert > 64)


def _desired_ksplit(token: int, expert: int) -> int:
    return 2 if _use_bypass(token, expert) else 0


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:
        try:
            bound = sig.bind_partial(*call_args, **call_kwargs)
            return (
                use_bypass,
                tuple((name, _cacheable(value)) for name, value in bound.arguments.items()),
            )
        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())),
    )


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


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

    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,
            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,
        )

    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,
        )
    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,
        )


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()
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,
    )
scrolls · 746 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 711471.

⋯ 19 unchanged lines
_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
⋯ 6 unchanged lines
(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),
+ }
+ )
def _replace_meta(meta, *, block_m=None, ksplit=None, stage1=None, stage2=None):
⋯ 132 unchanged lines
return None
+ 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
+
+
def _use_bypass(token: int, expert: int) -> bool:
return token <= 16 or (token <= 128 and expert > 64)
⋯ 40 unchanged lines
def _install_patch() -> None:
- global _PATCH_DONE
+ global _PATCH_DONE, _DIRECT_FUSED_MOE, _FUSED_MOE_MODULE
if _PATCH_DONE:
return
_PATCH_DONE = True
⋯ 1 unchanged lines
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")
⋯ 12 unchanged lines
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
⋯ 71 unchanged lines
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
⋯ 73 unchanged lines
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)
⋯ 4 unchanged lines
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,
+ 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,
+ )
+
+ 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,
+ )
+ 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,
+ )
+
+
+ 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()
def custom_kernel(data: input_t) -> output_t:
(
⋯ 17 unchanged lines
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(
+ 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,
- 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,
+ gate_up_weight_scale_shuffled,
+ down_weight_scale_shuffled,
+ hidden_pad,
+ intermediate_pad,
+ block_size_m,
)
scrolls · 393 diff lines total

Best evidence level for this revision: reported

JSON