diff --git a/tests/kernels/moe/test_cutlass_moe.py b/tests/kernels/moe/test_cutlass_moe.py index fa4351de7e2..274658ec0e8 100644 --- a/tests/kernels/moe/test_cutlass_moe.py +++ b/tests/kernels/moe/test_cutlass_moe.py @@ -11,8 +11,15 @@ import vllm.model_executor.layers.fused_moe.modular_kernel as mk from tests.kernels.moe.utils import make_dummy_moe_config from vllm import _custom_ops as ops from vllm.config import ParallelConfig, VllmConfig, set_current_vllm_config -from vllm.model_executor.layers.fused_moe import fused_experts, fused_topk -from vllm.model_executor.layers.fused_moe.activation import MoEActivation +from vllm.model_executor.layers.fused_moe import ( + ApplyMoEActivationConfig, + fused_experts, + fused_topk, +) +from vllm.model_executor.layers.fused_moe.activation import ( + MoEActivation, + apply_moe_activation_supported, +) from vllm.model_executor.layers.fused_moe.all2all_utils import ( maybe_make_prepare_finalize, ) @@ -24,8 +31,11 @@ from vllm.model_executor.layers.fused_moe.config import ( from vllm.model_executor.layers.fused_moe.experts.cutlass_moe import ( CutlassExpertsFp4, CutlassExpertsFp8, + CutlassExpertsMxfp4, + CutlassExpertsW4A8Fp8, run_cutlass_moe_fp8, ) +from vllm.model_executor.layers.fused_moe.oracle import nvfp4 as nvfp4_oracle from vllm.model_executor.layers.fused_moe.utils import moe_kernel_quantize_input from vllm.platforms import current_platform from vllm.utils.torch_utils import set_random_seed @@ -53,10 +63,72 @@ MNK_FACTORS = [ vllm_config = VllmConfig(parallel_config=ParallelConfig(pipeline_parallel_size=1)) -def test_cutlass_moe_supports_gelu_tanh_activation_metadata(): - assert CutlassExpertsFp8._supports_activation(MoEActivation.GELU_TANH) - assert CutlassExpertsFp4._supports_activation(MoEActivation.GELU_TANH) - assert CutlassExpertsFp4._supports_activation(MoEActivation.GELU_TANH_NO_MUL) +@pytest.mark.parametrize( + "experts_cls", + [ + CutlassExpertsFp8, + CutlassExpertsFp4, + CutlassExpertsMxfp4, + CutlassExpertsW4A8Fp8, + ], +) +@pytest.mark.parametrize("activation", list(MoEActivation)) +def test_cutlass_moe_activation_metadata_tracks_shared_apply( + experts_cls: type[mk.FusedMoEExperts], activation: MoEActivation +): + supports_shape = experts_cls._supports_no_act_and_mul() or activation.is_gated + expected = supports_shape and apply_moe_activation_supported(activation) + + assert experts_cls._supports_activation(activation) == expected + + +def test_cutlass_moe_forwards_shared_activation_parameters(): + moe_config = make_dummy_moe_config() + moe_config.swiglu_limit = 7.0 + moe_config.swiglu_alpha = 1.5 + moe_config.swiglu_beta = 0.25 + moe_config.activation_situ_beta = 2.0 + moe_config.activation_situ_linear_beta = 3.0 + quant_config = FusedMoEQuantConfig.make(gemm1_alpha=1.75) + + assert ApplyMoEActivationConfig.from_configs( + moe_config, quant_config + ) == ApplyMoEActivationConfig( + clamp_limit=7.0, + alpha=1.75, + beta=0.25, + activation_situ_beta=2.0, + activation_situ_linear_beta=3.0, + ) + + +@pytest.mark.parametrize( + ("backend", "expected"), + [ + ("cutlass", nvfp4_oracle.NvFp4MoeBackend.VLLM_CUTLASS), + ("humming", nvfp4_oracle.NvFp4MoeBackend.HUMMING), + ], +) +def test_nvfp4_clamp_allows_shared_activation_backends( + monkeypatch, backend: str, expected: nvfp4_oracle.NvFp4MoeBackend +): + class SupportedExperts: + @staticmethod + def is_supported_config(*args, **kwargs): + return True, None + + monkeypatch.setattr( + nvfp4_oracle, "backend_to_kernel_cls", lambda backend: [SupportedExperts] + ) + moe_config = make_dummy_moe_config() + moe_config.moe_backend = backend + moe_config.swiglu_limit = 7.0 + + selected, _ = nvfp4_oracle.select_nvfp4_moe_backend( + moe_config, weight_key=None, activation_key=None + ) + + assert selected == expected @dataclasses.dataclass diff --git a/tests/kernels/moe/test_moe.py b/tests/kernels/moe/test_moe.py index cee8971ae96..c70ce0b4695 100644 --- a/tests/kernels/moe/test_moe.py +++ b/tests/kernels/moe/test_moe.py @@ -27,6 +27,11 @@ from vllm.model_executor.layers.fused_moe import ( MoEActivation, fused_topk, ) +from vllm.model_executor.layers.fused_moe.activation import ( + ApplyMoEActivationConfig, + apply_moe_activation, + apply_moe_activation_supported, +) from vllm.model_executor.layers.fused_moe.config import ( FUSED_MOE_UNQUANTIZED_CONFIG, int4_w4a16_moe_quant_config, @@ -996,6 +1001,23 @@ def test_fused_marlin_moe( per_act_token_quant=True, ) + def instance_activation( + activation: MoEActivation, + output: torch.Tensor, + input: torch.Tensor, + *, + topk_ids: torch.Tensor | None = None, + expert_map: torch.Tensor | None = None, + ) -> None: + apply_moe_activation( + activation, + output, + input, + activation_config=ApplyMoEActivationConfig(), + topk_ids=topk_ids, + expert_map=expert_map, + ) + marlin_output = fused_marlin_moe( a, w1_data.qweight, @@ -1021,6 +1043,7 @@ def test_fused_marlin_moe( input_dtype=a_dtype, quant_type_id=b_type.id, is_k_full=is_k_full, + activation_func=instance_activation, ) torch.testing.assert_close(marlin_output, torch_output, atol=4e-2, rtol=0) @@ -1249,6 +1272,51 @@ def _make_humming_indexed_experts(activation: MoEActivation): return experts +@pytest.mark.parametrize("activation", list(MoEActivation)) +def test_humming_activation_metadata_tracks_shared_apply(activation: MoEActivation): + from vllm.model_executor.layers.fused_moe.experts.fused_humming_moe import ( + HummingExpertsBase, + ) + + assert HummingExpertsBase._supports_activation( + activation + ) == apply_moe_activation_supported(activation) + + +def test_humming_delegates_to_instance_activation(): + from types import SimpleNamespace + from unittest.mock import Mock + + from vllm.model_executor.layers.fused_moe.experts.fused_humming_moe import ( + HummingExpertsBase, + ) + + activation_func = Mock() + activation_config = ApplyMoEActivationConfig( + clamp_limit=7.0, + alpha=1.5, + beta=0.25, + activation_situ_beta=2.0, + activation_situ_linear_beta=3.0, + ) + experts = SimpleNamespace( + activation=activation_func, + activation_config=activation_config, + ) + input = torch.empty(1, 2) + output = torch.empty(1, 1) + + HummingExpertsBase.apply_activation( + experts, MoEActivation.SWIGLUOAI_UNINTERLEAVE, output, input + ) + + activation_func.assert_called_once_with( + activation=MoEActivation.SWIGLUOAI_UNINTERLEAVE, + input=input, + output=output, + ) + + @pytest.mark.parametrize( "activation", [ diff --git a/tests/kernels/moe/test_triton_moe_no_act_mul.py b/tests/kernels/moe/test_triton_moe_no_act_mul.py index 9d16ae5b63d..e9f6f51c220 100644 --- a/tests/kernels/moe/test_triton_moe_no_act_mul.py +++ b/tests/kernels/moe/test_triton_moe_no_act_mul.py @@ -11,7 +11,12 @@ import pytest import torch from tests.kernels.moe.utils import make_dummy_moe_config -from vllm.model_executor.layers.fused_moe.activation import MoEActivation +from vllm.model_executor.layers.fused_moe.activation import ( + ApplyMoEActivationConfig, + MoEActivation, + apply_moe_activation, + apply_moe_activation_supported, +) from vllm.model_executor.layers.fused_moe.config import ( FUSED_MOE_UNQUANTIZED_CONFIG, ) @@ -31,6 +36,49 @@ NO_MUL_ACTIVATIONS = [ ] +def test_apply_moe_activation_supported_contract(): + supported = { + activation + for activation in MoEActivation + if apply_moe_activation_supported(activation) + } + + assert supported == set(MoEActivation) - {MoEActivation.RELU2} + + +@pytest.mark.skipif(not current_platform.is_cuda(), reason="Requires CUDA") +@pytest.mark.parametrize( + "activation", + [ + activation + for activation in MoEActivation + if apply_moe_activation_supported(activation) + ], +) +@torch.inference_mode() +def test_supported_apply_moe_activation_executes(activation: MoEActivation): + output_width = 64 + input_width = output_width * 2 if activation.is_gated else output_width + input = torch.randn(2, input_width, device="cuda", dtype=torch.bfloat16) + output = torch.empty(2, output_width, device="cuda", dtype=torch.bfloat16) + activation_config = ApplyMoEActivationConfig() + if activation == MoEActivation.SITU: + activation_config = ApplyMoEActivationConfig(activation_situ_beta=1.0) + elif activation == MoEActivation.SWIGLUOAI_UNINTERLEAVE: + activation_config = ApplyMoEActivationConfig(clamp_limit=7.0) + + apply_moe_activation(activation, output, input, activation_config=activation_config) + + assert torch.isfinite(output).all() + + +@pytest.mark.parametrize("activation", list(MoEActivation)) +def test_triton_activation_metadata_tracks_shared_apply(activation: MoEActivation): + assert TritonExperts._supports_activation( + activation + ) == apply_moe_activation_supported(activation) + + def make_test_tensors( m: int, n: int, diff --git a/vllm/model_executor/layers/fused_moe/__init__.py b/vllm/model_executor/layers/fused_moe/__init__.py index 17f287556b7..120fbd15266 100644 --- a/vllm/model_executor/layers/fused_moe/__init__.py +++ b/vllm/model_executor/layers/fused_moe/__init__.py @@ -5,6 +5,7 @@ from contextlib import contextmanager from typing import Any from vllm.model_executor.layers.fused_moe.activation import ( + ApplyMoEActivationConfig, MoEActivation, activation_without_mul, apply_moe_activation, @@ -63,6 +64,7 @@ def get_config() -> dict[str, Any] | None: __all__ = [ + "ApplyMoEActivationConfig", "FusedMoEFactory", "FusedMoERouter", "FusedMoEConfig", diff --git a/vllm/model_executor/layers/fused_moe/activation.py b/vllm/model_executor/layers/fused_moe/activation.py index baa6136485e..6585bbe4f42 100644 --- a/vllm/model_executor/layers/fused_moe/activation.py +++ b/vllm/model_executor/layers/fused_moe/activation.py @@ -2,11 +2,19 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project """MoE activation function enum and utilities.""" +from dataclasses import dataclass from enum import Enum +from typing import TYPE_CHECKING import torch import torch.nn.functional as F +if TYPE_CHECKING: + from vllm.model_executor.layers.fused_moe.config import ( + FusedMoEConfig, + FusedMoEQuantConfig, + ) + class MoEActivation(Enum): """Activation functions for MoE layers.""" @@ -109,6 +117,66 @@ def activation_without_mul(activation: str) -> str: return MoEActivation.from_str(activation).without_mul().value +_APPLY_MOE_ACTIVATIONS = frozenset( + { + MoEActivation.SILU, + MoEActivation.GELU, + MoEActivation.GELU_TANH, + MoEActivation.SITU, + MoEActivation.SWIGLUOAI, + MoEActivation.SWIGLUOAI_UNINTERLEAVE, + MoEActivation.SWIGLUSTEP, + MoEActivation.SILU_NO_MUL, + MoEActivation.GELU_NO_MUL, + MoEActivation.GELU_TANH_NO_MUL, + MoEActivation.RELU2_NO_MUL, + } +) + + +def apply_moe_activation_supported(activation: MoEActivation) -> bool: + """Whether ``apply_moe_activation`` supports an activation.""" + return activation in _APPLY_MOE_ACTIVATIONS + + +@dataclass(frozen=True) +class ApplyMoEActivationConfig: + """Configuration forwarded to ``apply_moe_activation``.""" + + clamp_limit: float | None = None + alpha: float = 1.0 + beta: float = 0.0 + activation_situ_beta: float | None = None + activation_situ_linear_beta: float | None = None + + @classmethod + def from_configs( + cls, + moe_config: "FusedMoEConfig", + quant_config: "FusedMoEQuantConfig", + ) -> "ApplyMoEActivationConfig": + """Build from the model and quantization configurations.""" + clamp_limit = quant_config.gemm1_clamp_limit + if clamp_limit is None: + clamp_limit = moe_config.swiglu_limit + alpha = quant_config.gemm1_alpha + if alpha is None: + alpha = moe_config.swiglu_alpha + beta = quant_config.gemm1_beta + if beta is None: + beta = moe_config.swiglu_beta + return cls( + clamp_limit=clamp_limit, + alpha=1.0 if alpha is None else alpha, + beta=0.0 if beta is None else beta, + activation_situ_beta=moe_config.activation_situ_beta, + activation_situ_linear_beta=moe_config.activation_situ_linear_beta, + ) + + +_DEFAULT_APPLY_MOE_ACTIVATION_CONFIG = ApplyMoEActivationConfig() + + def silu_and_mul_with_clamp( output: torch.Tensor, input: torch.Tensor, @@ -130,20 +198,21 @@ def apply_moe_activation( output: torch.Tensor, input: torch.Tensor, *, - clamp_limit: float | None = None, - alpha: float = 1.0, - beta: float = 0.0, + activation_config: ApplyMoEActivationConfig | None = None, topk_ids: torch.Tensor | None = None, expert_map: torch.Tensor | None = None, - activation_situ_beta: float | None = None, - activation_situ_linear_beta: float | None = None, ) -> torch.Tensor: """Apply MoE activation function. - ``clamp_limit``/``alpha``/``beta`` (from the quant config) drive the clamped - SwiGLU kernels: ``SILU`` + ``clamp_limit`` and ``SWIGLUOAI_UNINTERLEAVE`` both - map to ``silu_and_mul_with_clamp``. Other activations ignore them. + The configuration drives specialized activation behavior. Routing tensors + remain per-call inputs because they depend on the current token assignment. """ + config = ( + _DEFAULT_APPLY_MOE_ACTIVATION_CONFIG + if activation_config is None + else activation_config + ) + assert input.dim() == 2, "Input must be 2D" assert output.dim() == 2, "Output must be 2D" if activation.is_gated: @@ -159,8 +228,10 @@ def apply_moe_activation( # Activations with gated multiplication (gate × activation(up)) if activation == MoEActivation.SILU: - if clamp_limit is not None: - silu_and_mul_with_clamp(output, input, clamp_limit, topk_ids, expert_map) + if config.clamp_limit is not None: + silu_and_mul_with_clamp( + output, input, config.clamp_limit, topk_ids, expert_map + ) else: torch.ops._C.silu_and_mul(output, input) elif activation == MoEActivation.GELU: @@ -175,23 +246,27 @@ def apply_moe_activation( # bypassed the config plumbing, so fail rather than silently use 1.0. # linear_beta is genuinely optional: <= 0 signals "unset" to the kernel # (up passed through), matching SituAndMul(linear_beta=None). - assert activation_situ_beta is not None, ( + assert config.activation_situ_beta is not None, ( "SITU requires activation_situ_beta from FusedMoEConfig" ) torch.ops._C.situ_and_mul( output, input, - activation_situ_beta, + config.activation_situ_beta, -1.0 - if activation_situ_linear_beta is None - else activation_situ_linear_beta, + if config.activation_situ_linear_beta is None + else config.activation_situ_linear_beta, ) elif activation == MoEActivation.SWIGLUOAI: torch.ops._C.swigluoai_and_mul(output, input) elif activation == MoEActivation.SWIGLUOAI_UNINTERLEAVE: # SwiGLU-OAI on packed w13 (gate = first half, up = second half). - assert clamp_limit is not None, "SWIGLUOAI_UNINTERLEAVE requires clamp_limit" - torch.ops._C.silu_and_mul_with_clamp(output, input, clamp_limit, alpha, beta) + assert config.clamp_limit is not None, ( + "SWIGLUOAI_UNINTERLEAVE requires clamp_limit" + ) + torch.ops._C.silu_and_mul_with_clamp( + output, input, config.clamp_limit, config.alpha, config.beta + ) elif activation == MoEActivation.SWIGLUSTEP: from vllm.model_executor.layers.activation import swiglustep_and_mul_triton diff --git a/vllm/model_executor/layers/fused_moe/experts/cutlass_moe.py b/vllm/model_executor/layers/fused_moe/experts/cutlass_moe.py index ce787fc09e0..80c6926fd4a 100644 --- a/vllm/model_executor/layers/fused_moe/experts/cutlass_moe.py +++ b/vllm/model_executor/layers/fused_moe/experts/cutlass_moe.py @@ -8,8 +8,10 @@ import vllm.model_executor.layers.fused_moe.modular_kernel as mk from vllm import _custom_ops as ops from vllm.logger import init_logger from vllm.model_executor.layers.fused_moe.activation import ( + ApplyMoEActivationConfig, MoEActivation, apply_moe_activation, + apply_moe_activation_supported, ) from vllm.model_executor.layers.fused_moe.config import ( FusedMoEConfig, @@ -76,6 +78,8 @@ def run_cutlass_moe_fp8( use_batched_format: bool, topk_weights: torch.Tensor | None, permute_scratch: MoEPermuteScratch | None, + *, + activation_config: ApplyMoEActivationConfig | None = None, ): a1q = hidden_states @@ -231,7 +235,12 @@ def run_cutlass_moe_fp8( per_out_ch, ) - apply_moe_activation(activation, act_out, mm1_out) + apply_moe_activation( + activation, + act_out, + mm1_out, + activation_config=activation_config, + ) a2q, a2q_scale = ops.scaled_fp8_quant( act_out, a2_scale, use_per_token_if_dynamic=per_act_token, output=quant_out @@ -319,12 +328,7 @@ class CutlassExpertsFp8Base(mk.FusedMoEExpertsModular): @staticmethod def _supports_activation(activation: MoEActivation) -> bool: - return activation in [ - MoEActivation.SILU, - MoEActivation.GELU, - MoEActivation.GELU_TANH, - MoEActivation.SWIGLUOAI, - ] + return activation.is_gated and apply_moe_activation_supported(activation) def finalize_weight_and_reduce_impl(self) -> mk.TopKWeightAndReduce: # Let PrepareAndFinalize::finalize() decide the impl. @@ -397,6 +401,7 @@ class CutlassExpertsFp8Base(mk.FusedMoEExpertsModular): use_batched_format, topk_weights, self._get_permute_scratch(), + activation_config=self.activation_config, ) @@ -512,6 +517,8 @@ def run_cutlass_moe_fp4( e: int, device: torch.device, apply_router_weight_on_input: bool = False, + *, + activation_config: ApplyMoEActivationConfig | None = None, ) -> None: """ MoE implementation for FP4 Inputs @@ -633,7 +640,9 @@ def run_cutlass_moe_fp4( blockscale_offsets[:-1], ) del rep_a_fp4, rep_a_blockscale - if activation == MoEActivation.SILU: + if activation == MoEActivation.SILU and ( + activation_config is None or activation_config.clamp_limit is None + ): # Fused SiLU+Mul+NVFP4 quantization # Note: c2 workspace is no longer needed since SiLU is fused with quantization. # c3 reuses workspace13 after c1 is consumed. @@ -641,7 +650,12 @@ def run_cutlass_moe_fp4( c1, a2_gscale, expert_offsets, blockscale_offsets, num_topk ) else: - apply_moe_activation(activation, c2, c1) + apply_moe_activation( + activation, + c2, + c1, + activation_config=activation_config, + ) int_fp4, int_blockscale = ops.scaled_fp4_experts_quant( c2, a2_gscale, expert_offsets, blockscale_offsets, num_topk ) @@ -716,17 +730,7 @@ class CutlassExpertsFp4(mk.FusedMoEExpertsModular): # fallback + separate fp4 quantization in run_cutlass_moe_fp4(). # Non-gated activations (_NO_MUL) are also supported for models # like Nemotron-Nano that don't use gated MLP. - return activation in [ - MoEActivation.SILU, - MoEActivation.GELU, - MoEActivation.GELU_TANH, - MoEActivation.SWIGLUOAI, - MoEActivation.SWIGLUSTEP, - MoEActivation.SILU_NO_MUL, - MoEActivation.GELU_NO_MUL, - MoEActivation.GELU_TANH_NO_MUL, - MoEActivation.RELU2_NO_MUL, - ] + return apply_moe_activation_supported(activation) @staticmethod def _supports_parallel_config(moe_parallel_config: FusedMoEParallelConfig) -> bool: @@ -807,6 +811,7 @@ class CutlassExpertsFp4(mk.FusedMoEExpertsModular): e=e, device=hidden_states.device, apply_router_weight_on_input=apply_router_weight_on_input, + activation_config=self.activation_config, ) @@ -828,6 +833,8 @@ def run_cutlass_moe_mxfp4( e: int, device: torch.device, apply_router_weight_on_input: bool = False, + *, + activation_config: ApplyMoEActivationConfig | None = None, ) -> None: """MXFP4 x MXFP4 MoE implementation using CUTLASS grouped GEMM.""" is_gated = activation.is_gated @@ -909,12 +916,19 @@ def run_cutlass_moe_mxfp4( blockscale_offsets[:-1], ) del rep_a_fp4, rep_a_blockscale - if activation == MoEActivation.SILU: + if activation == MoEActivation.SILU and ( + activation_config is None or activation_config.clamp_limit is None + ): int_fp4, int_blockscale = ops.silu_and_mul_mxfp4_experts_quant( c1, expert_offsets, blockscale_offsets, e, num_topk ) else: - apply_moe_activation(activation, c2, c1) + apply_moe_activation( + activation, + c2, + c1, + activation_config=activation_config, + ) int_fp4, int_blockscale = ops.mxfp4_experts_quant( c2, expert_offsets, blockscale_offsets, e, num_topk ) @@ -1021,15 +1035,7 @@ class CutlassExpertsMxfp4(mk.FusedMoEExpertsModular): @staticmethod def _supports_activation(activation: MoEActivation) -> bool: - return activation in [ - MoEActivation.SILU, - MoEActivation.GELU, - MoEActivation.SWIGLUOAI, - MoEActivation.SWIGLUSTEP, - MoEActivation.SILU_NO_MUL, - MoEActivation.GELU_NO_MUL, - MoEActivation.RELU2_NO_MUL, - ] + return apply_moe_activation_supported(activation) @staticmethod def _supports_parallel_config( @@ -1102,6 +1108,7 @@ class CutlassExpertsMxfp4(mk.FusedMoEExpertsModular): e=e, device=hidden_states.device, apply_router_weight_on_input=apply_router_weight_on_input, + activation_config=self.activation_config, ) @@ -1139,6 +1146,8 @@ def run_cutlass_moe_w4a8_fp8( topk_weights: torch.Tensor | None, group_size: int, permute_scratch: MoEPermuteScratch | None, + *, + activation_config: ApplyMoEActivationConfig | None = None, ): a1q = hidden_states M = a1q.size(0) @@ -1218,7 +1227,12 @@ def run_cutlass_moe_w4a8_fp8( s_strides1, ) - apply_moe_activation(activation, act_out, mm1_out) + apply_moe_activation( + activation, + act_out, + mm1_out, + activation_config=activation_config, + ) a2q, a2q_scale = ops.scaled_fp8_quant( act_out, a2_scale, use_per_token_if_dynamic=per_act_token, output=quant_out @@ -1330,11 +1344,7 @@ class CutlassExpertsW4A8Fp8(mk.FusedMoEExpertsModular): @staticmethod def _supports_activation(activation: MoEActivation) -> bool: - return activation in ( - MoEActivation.SILU, - MoEActivation.GELU, - MoEActivation.SWIGLUOAI, - ) + return activation.is_gated and apply_moe_activation_supported(activation) @staticmethod def _supports_parallel_config(moe_parallel_config: FusedMoEParallelConfig) -> bool: @@ -1438,4 +1448,5 @@ class CutlassExpertsW4A8Fp8(mk.FusedMoEExpertsModular): topk_weights, self.group_size, self._get_permute_scratch(), + activation_config=self.activation_config, ) diff --git a/vllm/model_executor/layers/fused_moe/experts/fused_humming_moe.py b/vllm/model_executor/layers/fused_moe/experts/fused_humming_moe.py index 2f073fffccb..e0ae629d1b4 100644 --- a/vllm/model_executor/layers/fused_moe/experts/fused_humming_moe.py +++ b/vllm/model_executor/layers/fused_moe/experts/fused_humming_moe.py @@ -12,7 +12,10 @@ import vllm.model_executor.layers.fused_moe.modular_kernel as mk from vllm import envs from vllm.forward_context import get_forward_context from vllm.logger import init_logger -from vllm.model_executor.layers.fused_moe.activation import MoEActivation +from vllm.model_executor.layers.fused_moe.activation import ( + MoEActivation, + apply_moe_activation_supported, +) from vllm.model_executor.layers.fused_moe.config import ( FusedMoEConfig, FusedMoEParallelConfig, @@ -246,18 +249,7 @@ class HummingExpertsBase(mk.FusedMoEExpertsModular): def _supports_activation(activation: MoEActivation) -> bool: # Humming uses apply_moe_activation() callback for activation, # so any activation supported there can be used here. - return activation in [ - MoEActivation.SILU, - MoEActivation.GELU, - MoEActivation.GELU_TANH, - MoEActivation.SWIGLUOAI, - MoEActivation.SITU, - MoEActivation.SWIGLUSTEP, - MoEActivation.SILU_NO_MUL, - MoEActivation.GELU_NO_MUL, - MoEActivation.GELU_TANH_NO_MUL, - MoEActivation.RELU2_NO_MUL, - ] + return apply_moe_activation_supported(activation) @staticmethod def _supports_parallel_config(moe_parallel_config: FusedMoEParallelConfig) -> bool: @@ -510,11 +502,22 @@ class HummingExpertsBase(mk.FusedMoEExpertsModular): output: torch.Tensor, input: torch.Tensor, ) -> None: - swiglu_limit = self.quant_config.gemm1_clamp_limit - if activation == MoEActivation.SILU and swiglu_limit is not None: - swiglu_limit_func(output=output, input=input, swiglu_limit=swiglu_limit) + activation_config = self.activation_config + if ( + activation == MoEActivation.SILU + and activation_config.clamp_limit is not None + ): + swiglu_limit_func( + output=output, + input=input, + swiglu_limit=activation_config.clamp_limit, + ) else: - self.activation(activation=activation, input=input, output=output) + self.activation( + activation=activation, + input=input, + output=output, + ) class HummingIndexedExperts(HummingExpertsBase): diff --git a/vllm/model_executor/layers/fused_moe/experts/marlin_moe.py b/vllm/model_executor/layers/fused_moe/experts/marlin_moe.py index 1125130c890..c3eb5eefb36 100644 --- a/vllm/model_executor/layers/fused_moe/experts/marlin_moe.py +++ b/vllm/model_executor/layers/fused_moe/experts/marlin_moe.py @@ -10,8 +10,10 @@ import torch import vllm._custom_ops as ops import vllm.model_executor.layers.fused_moe.modular_kernel as mk from vllm.model_executor.layers.fused_moe.activation import ( + ApplyMoEActivationConfig, MoEActivation, apply_moe_activation, + apply_moe_activation_supported, ) from vllm.model_executor.layers.fused_moe.config import ( FusedMoEConfig, @@ -72,7 +74,7 @@ def _fused_marlin_moe( expert_ids: torch.Tensor, num_tokens_post_padded: torch.Tensor, activation: MoEActivation = MoEActivation.SILU, - activation_func: Callable[..., None] = apply_moe_activation, + activation_func: Callable[..., None] | None = None, topk_ids: torch.Tensor | None = None, input_global_scale1: torch.Tensor | None = None, input_global_scale2: torch.Tensor | None = None, @@ -90,11 +92,7 @@ def _fused_marlin_moe( output: torch.Tensor | None = None, input_dtype: torch.dtype | None = None, is_k_full: bool = True, - clamp_limit: float | None = None, - gemm1_alpha: float = 1.0, - gemm1_beta: float = 0.0, - activation_situ_beta: float | None = None, - activation_situ_linear_beta: float | None = None, + activation_config: ApplyMoEActivationConfig | None = None, ) -> torch.Tensor: assert hidden_states.ndim == 2 M, K = hidden_states.size() @@ -162,20 +160,29 @@ def _fused_marlin_moe( use_fp32_reduce=True, is_zp_float=False, ) - # apply_moe_activation fuses the clamp/gate params: SILU + clamp_limit and - # SWIGLUOAI_UNINTERLEAVE both map to the silu_and_mul_with_clamp kernel. - activation_func( - activation, - intermediate_cache2, - intermediate_cache1.view(-1, w13_num_shards * N), - clamp_limit=clamp_limit, - alpha=gemm1_alpha, - beta=gemm1_beta, - topk_ids=topk_ids, - expert_map=expert_map, - activation_situ_beta=activation_situ_beta, - activation_situ_linear_beta=activation_situ_linear_beta, - ) + activation_input = intermediate_cache1.view(-1, w13_num_shards * N) + if activation_func is None: + config = ( + ApplyMoEActivationConfig() + if activation_config is None + else activation_config + ) + apply_moe_activation( + activation, + intermediate_cache2, + activation_input, + activation_config=config, + topk_ids=topk_ids, + expert_map=expert_map, + ) + else: + activation_func( + activation, + intermediate_cache2, + activation_input, + topk_ids=topk_ids, + expert_map=expert_map, + ) if output is None: output = intermediate_cache3 @@ -238,7 +245,7 @@ def fused_marlin_moe( apply_router_weight_on_input: bool = False, global_num_experts: int = -1, activation: MoEActivation = MoEActivation.SILU, - activation_func: Callable[..., None] = apply_moe_activation, + activation_func: Callable[..., None] | None = None, moe_sum: Callable[..., torch.Tensor | None] | None = None, expert_map: torch.Tensor | None = None, input_global_scale1: torch.Tensor | None = None, @@ -257,11 +264,7 @@ def fused_marlin_moe( is_k_full: bool = True, output: torch.Tensor | None = None, input_dtype: torch.dtype | None = None, - clamp_limit: float | None = None, - gemm1_alpha: float = 1.0, - gemm1_beta: float = 0.0, - activation_situ_beta: float | None = None, - activation_situ_linear_beta: float | None = None, + activation_config: ApplyMoEActivationConfig | None = None, ) -> torch.Tensor: """ This function computes a Mixture of Experts (MoE) layer using two sets of @@ -363,8 +366,7 @@ def fused_marlin_moe( num_tokens_post_padded=num_tokens_post_padded, activation=activation, activation_func=activation_func, - activation_situ_beta=activation_situ_beta, - activation_situ_linear_beta=activation_situ_linear_beta, + activation_config=activation_config, input_global_scale1=input_global_scale1, input_global_scale2=input_global_scale2, global_scale1=global_scale1, @@ -381,9 +383,6 @@ def fused_marlin_moe( output=None, input_dtype=input_dtype, is_k_full=is_k_full, - clamp_limit=clamp_limit, - gemm1_alpha=gemm1_alpha, - gemm1_beta=gemm1_beta, ).view(-1, topk, K) if output is None: @@ -428,12 +427,8 @@ def batched_fused_marlin_moe( is_k_full: bool = True, output: torch.Tensor | None = None, input_dtype: torch.dtype | None = None, - clamp_limit: float | None = None, - gemm1_alpha: float = 1.0, - gemm1_beta: float = 0.0, - activation_func: Callable[..., None] = apply_moe_activation, - activation_situ_beta: float | None = None, - activation_situ_linear_beta: float | None = None, + activation_func: Callable[..., None] | None = None, + activation_config: ApplyMoEActivationConfig | None = None, ) -> torch.Tensor: """ This function massages the inputs so the batched hidden_states can be @@ -542,8 +537,7 @@ def batched_fused_marlin_moe( apply_router_weight_on_input=apply_router_weight_on_input, activation=activation, activation_func=activation_func, - activation_situ_beta=activation_situ_beta, - activation_situ_linear_beta=activation_situ_linear_beta, + activation_config=activation_config, expert_map=expert_map, block_size_m=block_size_m, sorted_token_ids=sorted_token_ids, @@ -565,9 +559,6 @@ def batched_fused_marlin_moe( output=output.view(-1, K) if output is not None else output, input_dtype=input_dtype, is_k_full=is_k_full, - clamp_limit=clamp_limit, - gemm1_alpha=gemm1_alpha, - gemm1_beta=gemm1_beta, ) output = output.view(B, BATCH_TOKENS_MAX, K) @@ -602,16 +593,6 @@ class MarlinExpertsBase(mk.FusedMoEExpertsModular): self.w2_g_idx_sort_indices = w2_g_idx_sort_indices self.is_k_full = is_k_full self.input_dtype = get_marlin_input_dtype() - self.gemm1_clamp_limit = quant_config.gemm1_clamp_limit - # Gated-activation params (used by SWIGLUOAI_UNINTERLEAVE on packed w13). - # silu == swigluoai with alpha=1, beta=0; configs that don't set these - # (plain silu) fall back to the silu identity. - self.gemm1_alpha = ( - quant_config.gemm1_alpha if quant_config.gemm1_alpha is not None else 1.0 - ) - self.gemm1_beta = ( - quant_config.gemm1_beta if quant_config.gemm1_beta is not None else 0.0 - ) super().__init__( moe_config=moe_config, @@ -655,19 +636,7 @@ class MarlinExpertsBase(mk.FusedMoEExpertsModular): def _supports_activation(activation: MoEActivation) -> bool: # Marlin uses apply_moe_activation() callback for activation, # so any activation supported there can be used here. - return activation in [ - MoEActivation.SILU, - MoEActivation.GELU, - MoEActivation.GELU_TANH, - MoEActivation.SITU, - MoEActivation.SWIGLUOAI, - MoEActivation.SWIGLUOAI_UNINTERLEAVE, - MoEActivation.SWIGLUSTEP, - MoEActivation.SILU_NO_MUL, - MoEActivation.GELU_NO_MUL, - MoEActivation.GELU_TANH_NO_MUL, - MoEActivation.RELU2_NO_MUL, - ] + return apply_moe_activation_supported(activation) @staticmethod def _supports_parallel_config(moe_parallel_config: FusedMoEParallelConfig) -> bool: @@ -808,10 +777,7 @@ class MarlinExperts(LoRAExpertsMixin, MarlinExpertsBase): global_num_experts=global_num_experts, activation=activation, activation_func=self.activation, - activation_situ_beta=self.moe_config.activation_situ_beta, - activation_situ_linear_beta=( - self.moe_config.activation_situ_linear_beta - ), + activation_config=self.activation_config, moe_sum=self.moe_sum, expert_map=expert_map, output=output, @@ -825,9 +791,6 @@ class MarlinExperts(LoRAExpertsMixin, MarlinExpertsBase): sort_indices2=self.w2_g_idx_sort_indices, is_k_full=self.is_k_full, input_dtype=self.input_dtype, - clamp_limit=self.gemm1_clamp_limit, - gemm1_alpha=self.gemm1_alpha, - gemm1_beta=self.gemm1_beta, ) return @@ -847,13 +810,8 @@ class MarlinExperts(LoRAExpertsMixin, MarlinExpertsBase): act_output: torch.Tensor, act_input: torch.Tensor, *, - clamp_limit: float | None = None, - alpha: float = 1.0, - beta: float = 0.0, topk_ids: torch.Tensor | None = None, expert_map: torch.Tensor | None = None, - activation_situ_beta: float | None = None, - activation_situ_linear_beta: float | None = None, ) -> None: # act_input = intermediate_cache1 (M*topk, 2N for gated) # act_output = intermediate_cache2 (M*topk, N) @@ -887,13 +845,8 @@ class MarlinExperts(LoRAExpertsMixin, MarlinExpertsBase): act_enum, act_output, act_input, - clamp_limit=clamp_limit, - alpha=alpha, - beta=beta, topk_ids=topk_ids, expert_map=expert_map, - activation_situ_beta=activation_situ_beta, - activation_situ_linear_beta=activation_situ_linear_beta, ) lora_state["cache2"] = act_output @@ -941,8 +894,7 @@ class MarlinExperts(LoRAExpertsMixin, MarlinExpertsBase): global_num_experts=global_num_experts, activation=activation, activation_func=activation_with_lora, - activation_situ_beta=self.moe_config.activation_situ_beta, - activation_situ_linear_beta=self.moe_config.activation_situ_linear_beta, + activation_config=self.activation_config, moe_sum=moe_sum_with_lora, expert_map=expert_map, output=output, @@ -954,9 +906,6 @@ class MarlinExperts(LoRAExpertsMixin, MarlinExpertsBase): sort_indices2=self.w2_g_idx_sort_indices, is_k_full=self.is_k_full, input_dtype=self.input_dtype, - clamp_limit=self.gemm1_clamp_limit, - gemm1_alpha=self.gemm1_alpha, - gemm1_beta=self.gemm1_beta, ) def moe_sum( @@ -1051,9 +1000,6 @@ class BatchedMarlinExperts(MarlinExpertsBase): act: MoEActivation, act_output: torch.Tensor, act_input: torch.Tensor, - *, - activation_situ_beta: float | None = None, - activation_situ_linear_beta: float | None = None, **kwargs, ) -> None: if act != MoEActivation.SITU: @@ -1061,20 +1007,18 @@ class BatchedMarlinExperts(MarlinExpertsBase): act, act_output, act_input, - activation_situ_beta=activation_situ_beta, - activation_situ_linear_beta=activation_situ_linear_beta, **kwargs, ) return num_experts, max_num_tokens = hidden_states.shape[:2] - beta = 1.0 if activation_situ_beta is None else activation_situ_beta - linear_beta = activation_situ_linear_beta + beta = self.activation_config.activation_situ_beta + linear_beta = self.activation_config.activation_situ_linear_beta torch.ops._C.masked_situ_and_mul( act_output.view(num_experts, max_num_tokens, -1), act_input.view(num_experts, max_num_tokens, -1), expert_tokens_meta.expert_num_tokens, - beta, + 1.0 if beta is None else beta, -1.0 if linear_beta is None else linear_beta, ) @@ -1105,10 +1049,6 @@ class BatchedMarlinExperts(MarlinExpertsBase): w2_zeros=self.w2_zp, input_dtype=self.input_dtype, is_k_full=self.is_k_full, - clamp_limit=self.gemm1_clamp_limit, - gemm1_alpha=self.gemm1_alpha, - gemm1_beta=self.gemm1_beta, activation_func=activation_func, - activation_situ_beta=self.moe_config.activation_situ_beta, - activation_situ_linear_beta=self.moe_config.activation_situ_linear_beta, + activation_config=self.activation_config, ) diff --git a/vllm/model_executor/layers/fused_moe/experts/mxfp8_emulation_moe.py b/vllm/model_executor/layers/fused_moe/experts/mxfp8_emulation_moe.py index ad6083251fb..d330d5f4dc6 100644 --- a/vllm/model_executor/layers/fused_moe/experts/mxfp8_emulation_moe.py +++ b/vllm/model_executor/layers/fused_moe/experts/mxfp8_emulation_moe.py @@ -104,16 +104,13 @@ class Mxfp8EmulationTritonExperts(Mxfp8TritonExpertsBase): ) if activation == MoEActivation.SWIGLUOAI_UNINTERLEAVE: - limit = self.quant_config.gemm1_clamp_limit - if limit is None: + if self.activation_config.clamp_limit is None: raise ValueError("SWIGLUOAI_UNINTERLEAVE requires gemm1_clamp_limit") apply_moe_activation( activation, output, input, - clamp_limit=float(limit), - alpha=self.gemm1_alpha, - beta=self.gemm1_beta, + activation_config=self.activation_config, ) return super().activation(activation, output, input) diff --git a/vllm/model_executor/layers/fused_moe/experts/mxfp8_native_moe.py b/vllm/model_executor/layers/fused_moe/experts/mxfp8_native_moe.py index e8c7dc96921..c8063ffb4d6 100644 --- a/vllm/model_executor/layers/fused_moe/experts/mxfp8_native_moe.py +++ b/vllm/model_executor/layers/fused_moe/experts/mxfp8_native_moe.py @@ -362,9 +362,7 @@ class Mxfp8NativeTritonExperts(Mxfp8TritonExpertsBase): expert_tokens_meta: mk.ExpertTokensMetadata | None, apply_router_weight_on_input: bool, ): - # `self.gemm1_alpha` and `self.gemm1_beta`` are set by `TritonExperts.__init__`. - limit = self.quant_config.gemm1_clamp_limit - limit = None if limit is None else float(limit) + activation_config = self.activation_config out = fused_moe_mxfp8_native( hidden_states, w1, @@ -373,9 +371,9 @@ class Mxfp8NativeTritonExperts(Mxfp8TritonExpertsBase): self.w2_scale_val, topk_weights, topk_ids, - alpha=self.gemm1_alpha, - beta=self.gemm1_beta, - limit=limit, + alpha=activation_config.alpha, + beta=activation_config.beta, + limit=activation_config.clamp_limit, global_num_experts=global_num_experts, expert_map=expert_map, ) diff --git a/vllm/model_executor/layers/fused_moe/experts/triton_moe.py b/vllm/model_executor/layers/fused_moe/experts/triton_moe.py index d1bab54dda1..e5a6ab65865 100644 --- a/vllm/model_executor/layers/fused_moe/experts/triton_moe.py +++ b/vllm/model_executor/layers/fused_moe/experts/triton_moe.py @@ -6,7 +6,10 @@ import torch import vllm.model_executor.layers.fused_moe.modular_kernel as mk from vllm import _custom_ops as ops -from vllm.model_executor.layers.fused_moe.activation import MoEActivation +from vllm.model_executor.layers.fused_moe.activation import ( + MoEActivation, + apply_moe_activation_supported, +) from vllm.model_executor.layers.fused_moe.config import ( FusedMoEConfig, FusedMoEParallelConfig, @@ -69,15 +72,6 @@ class TritonExperts(LoRAExpertsMixin, mk.FusedMoEExpertsModular): self.quantization_emulation = False super().__init__(moe_config, quant_config) - self.gemm1_clamp_limit = quant_config.gemm1_clamp_limit - # Gated-activation params: silu == swigluoai with alpha=1, beta=0. - self.gemm1_alpha = ( - quant_config.gemm1_alpha if quant_config.gemm1_alpha is not None else 1.0 - ) - self.gemm1_beta = ( - quant_config.gemm1_beta if quant_config.gemm1_beta is not None else 0.0 - ) - @staticmethod def activation_format() -> mk.FusedMoEActivationFormat: return mk.FusedMoEActivationFormat.Standard @@ -136,19 +130,7 @@ class TritonExperts(LoRAExpertsMixin, mk.FusedMoEExpertsModular): @staticmethod def _supports_activation(activation: MoEActivation) -> bool: - return activation in [ - MoEActivation.SILU, - MoEActivation.GELU, - MoEActivation.GELU_TANH, - MoEActivation.SITU, - MoEActivation.SWIGLUOAI, - MoEActivation.SWIGLUOAI_UNINTERLEAVE, - MoEActivation.SWIGLUSTEP, - MoEActivation.SILU_NO_MUL, - MoEActivation.GELU_NO_MUL, - MoEActivation.GELU_TANH_NO_MUL, - MoEActivation.RELU2_NO_MUL, - ] + return apply_moe_activation_supported(activation) @staticmethod def _supports_parallel_config(moe_parallel_config: FusedMoEParallelConfig) -> bool: @@ -171,17 +153,18 @@ class TritonExperts(LoRAExpertsMixin, mk.FusedMoEExpertsModular): input: torch.Tensor, **kwargs, ) -> None: - gemm1_clamp_limit = self.quant_config.gemm1_clamp_limit - if activation == MoEActivation.SILU and gemm1_clamp_limit is not None: - swiglu_limit_func(output, input, float(gemm1_clamp_limit)) + activation_config = self.activation_config + if ( + activation == MoEActivation.SILU + and activation_config.clamp_limit is not None + ): + swiglu_limit_func(output, input, activation_config.clamp_limit) return # SWIGLUOAI_UNINTERLEAVE routes to the silu_and_mul_with_clamp kernel and - # needs the clamped-SwiGLU params (gemm1_clamp_limit/alpha/beta read from - # the quant config in __init__) forwarded; without a clamp_limit it - # asserts. Other activations ignore alpha/beta/clamp_limit. + # requires a clamp limit. Other activations ignore these parameters. if activation == MoEActivation.SWIGLUOAI_UNINTERLEAVE: - assert gemm1_clamp_limit is not None, ( + assert activation_config.clamp_limit is not None, ( "SWIGLUOAI_UNINTERLEAVE requires gemm1_clamp_limit" ) @@ -189,9 +172,6 @@ class TritonExperts(LoRAExpertsMixin, mk.FusedMoEExpertsModular): activation, output, input, - clamp_limit=gemm1_clamp_limit, - alpha=self.gemm1_alpha, - beta=self.gemm1_beta, ) def workspace_shapes( diff --git a/vllm/model_executor/layers/fused_moe/modular_kernel.py b/vllm/model_executor/layers/fused_moe/modular_kernel.py index f7dc5dc2c05..7dba2d271a3 100644 --- a/vllm/model_executor/layers/fused_moe/modular_kernel.py +++ b/vllm/model_executor/layers/fused_moe/modular_kernel.py @@ -12,6 +12,7 @@ import torch import vllm.envs as envs from vllm.logger import init_logger from vllm.model_executor.layers.fused_moe.activation import ( + ApplyMoEActivationConfig, MoEActivation, apply_moe_activation, ) @@ -497,6 +498,9 @@ class FusedMoEExperts(ABC): self.moe_config = moe_config self.quant_config = quant_config + self.activation_config = ApplyMoEActivationConfig.from_configs( + moe_config, quant_config + ) self.max_num_tokens = max_num_tokens self.num_dispatchers = num_dispatchers @@ -885,33 +889,16 @@ class FusedMoEExpertsModular(FusedMoEExperts): output: torch.Tensor, input: torch.Tensor, *, - clamp_limit: float | None = None, - alpha: float = 1.0, - beta: float = 0.0, topk_ids: torch.Tensor | None = None, expert_map: torch.Tensor | None = None, - activation_situ_beta: float | None = None, - activation_situ_linear_beta: float | None = None, ) -> None: apply_moe_activation( activation, output, input, - clamp_limit=clamp_limit, - alpha=alpha, - beta=beta, + activation_config=self.activation_config, topk_ids=topk_ids, expert_map=expert_map, - activation_situ_beta=( - self.moe_config.activation_situ_beta - if activation_situ_beta is None - else activation_situ_beta - ), - activation_situ_linear_beta=( - self.moe_config.activation_situ_linear_beta - if activation_situ_linear_beta is None - else activation_situ_linear_beta - ), ) @abstractmethod diff --git a/vllm/model_executor/layers/fused_moe/oracle/nvfp4.py b/vllm/model_executor/layers/fused_moe/oracle/nvfp4.py index 7852b4db847..8439490d300 100644 --- a/vllm/model_executor/layers/fused_moe/oracle/nvfp4.py +++ b/vllm/model_executor/layers/fused_moe/oracle/nvfp4.py @@ -190,7 +190,9 @@ def select_nvfp4_moe_backend( NVFP4_BACKENDS_WITH_CLAMP = { NvFp4MoeBackend.FLASHINFER_TRTLLM, NvFp4MoeBackend.FLASHINFER_CUTLASS, + NvFp4MoeBackend.VLLM_CUTLASS, NvFp4MoeBackend.MARLIN, + NvFp4MoeBackend.HUMMING, } if config.swiglu_limit is not None: @@ -258,8 +260,9 @@ def select_nvfp4_moe_backend( raise ValueError( f"Model sets swiglu_limit={config.swiglu_limit}, but the " f"explicitly requested moe_backend={runner_backend!r} does " - f"not apply the SwiGLU clamp. Use 'flashinfer_trtllm' or " - f"'flashinfer_cutlass' instead." + f"not apply the SwiGLU clamp. Use 'flashinfer_trtllm', " + f"'flashinfer_cutlass', 'cutlass', 'marlin', or 'humming' " + f"instead." ) return _return_or_raise( requested_backend, config, weight_key, activation_key, activation_format