Skip to content

vllm.model_executor.layers.fused_moe.experts.flashinfer_cutedsl_moe ¶

Classes:

FlashInferCuteDSLExperts ¶

Bases: FusedMoEExpertsModular

CuteDSL NvFP4 MoE experts using the FlashInfer functional API.

Uses Standard activation format (non-batched). The kernel handles routing, expert computation, and reduction internally. Supports expert parallelism natively.

Source code in vllm/model_executor/layers/fused_moe/experts/flashinfer_cutedsl_moe.py
class FlashInferCuteDSLExperts(mk.FusedMoEExpertsModular):
    """CuteDSL NvFP4 MoE experts using the FlashInfer functional API.

    Uses Standard activation format (non-batched). The kernel handles
    routing, expert computation, and reduction internally.
    Supports expert parallelism natively.
    """

    def __init__(
        self,
        moe_config: FusedMoEConfig,
        quant_config: FusedMoEQuantConfig,
        per_token_activation: bool = False,
    ):
        super().__init__(
            moe_config=moe_config,
            quant_config=quant_config,
        )
        assert quant_config.quant_dtype == "nvfp4", (
            "Only nvfp4 quantization is currently supported."
        )
        self.out_dtype = moe_config.in_dtype
        self.hidden_dim = moe_config.hidden_dim
        self.intermediate_size_per_partition = (
            moe_config.intermediate_size_per_partition
        )
        self.topk = moe_config.experts_per_token
        self.local_num_experts = moe_config.num_local_experts
        self.global_num_experts = moe_config.num_experts
        self.ep_rank = moe_config.moe_parallel_config.ep_rank
        self.local_expert_offset = self.ep_rank * self.local_num_experts
        self.gemm1_alpha = quant_config.gemm1_alpha
        self.gemm1_beta = quant_config.gemm1_beta
        self.gemm1_clamp_limit = quant_config.gemm1_clamp_limit
        self.situ_beta = moe_config.activation_situ_beta
        self.situ_linear_beta = moe_config.activation_situ_linear_beta
        self.per_token_activation = per_token_activation
        self.per_token_global_scale = None
        if per_token_activation:
            assert quant_config.a2_gscale is not None
            self.per_token_global_scale = quant_config.a2_gscale.new_full(
                (1,),
                NVFP4_PER_TOKEN_BASE_GLOBAL_SCALE,
            )

    def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
        layer.w13_weight_scale_2.data.mul_(layer.w13_input_scale)
        layer.w2_weight_scale_2.data.mul_(layer.w2_input_scale)

    @staticmethod
    def activation_format() -> mk.FusedMoEActivationFormat:
        return mk.FusedMoEActivationFormat.Standard

    @staticmethod
    def _supports_current_device() -> bool:
        p = current_platform
        return (
            p.is_cuda()
            and p.is_device_capability_family(100)
            and has_flashinfer_cutedsl_moe_nvfp4()
        )

    @staticmethod
    def _supports_no_act_and_mul() -> bool:
        return True

    @staticmethod
    def _supports_quant_scheme(
        weight_key: QuantKey | None,
        activation_key: QuantKey | None,
    ) -> bool:
        SUPPORTED_W_A = [
            (kNvfp4Static, kNvfp4Dynamic),
            (kNvfp4Static, kNvfp4DynamicToken),
        ]
        return (weight_key, activation_key) in SUPPORTED_W_A

    @staticmethod
    def _supports_activation(activation: MoEActivation) -> bool:
        return activation in (
            MoEActivation.SILU,
            MoEActivation.SWIGLUOAI,
            MoEActivation.SWIGLUOAI_UNINTERLEAVE,
            MoEActivation.RELU2_NO_MUL,
            MoEActivation.SITU,
        )

    @staticmethod
    def _supports_parallel_config(
        moe_parallel_config: FusedMoEParallelConfig,
    ) -> bool:
        return True

    def finalize_weight_and_reduce_impl(self) -> mk.TopKWeightAndReduce:
        return TopKWeightAndReduceNoOP()

    @property
    def expects_unquantized_inputs(self) -> bool:
        return self.per_token_activation

    def workspace_shapes(
        self,
        M: int,
        N: int,
        K: int,
        topk: int,
        global_num_experts: int,
        local_num_experts: int,
        expert_tokens_meta: mk.ExpertTokensMetadata | None,
        activation: MoEActivation,
    ) -> tuple[tuple[int, ...], tuple[int, ...], tuple[int, ...]]:
        workspace1 = (0,)
        workspace2 = (0,)
        expected_hidden_dim = K if self.expects_unquantized_inputs else K * 2
        assert self.hidden_dim == expected_hidden_dim
        output = (M, self.hidden_dim)
        return (workspace1, workspace2, output)

    def apply(
        self,
        output: torch.Tensor,
        hidden_states: torch.Tensor,
        w1: torch.Tensor,
        w2: torch.Tensor,
        topk_weights: torch.Tensor,
        topk_ids: torch.Tensor,
        activation: MoEActivation,
        global_num_experts: int,
        expert_map: torch.Tensor | None,
        a1q_scale: torch.Tensor | None,
        a2_scale: torch.Tensor | None,
        workspace13: torch.Tensor | None,
        workspace2: torch.Tensor | None,
        expert_tokens_meta: mk.ExpertTokensMetadata | None,
        apply_router_weight_on_input: bool | None,
    ):
        assert self.quant_dtype == "nvfp4"
        assert self.w1_scale is not None
        assert self.w2_scale is not None

        if self.expects_unquantized_inputs:
            hidden_states, block_scale, per_token_scale = (
                quantize_nvfp4_per_token_input(hidden_states)
            )
            fc2_input_scale = self.per_token_global_scale
        else:
            assert a1q_scale is not None
            block_scale = a1q_scale
            per_token_scale = None
            fc2_input_scale = self.a2_gscale

        assert block_scale is not None
        assert fc2_input_scale is not None

        # a1q_scale is (M, K//16) float8_e4m3fn from fp4_quantize.
        # The functional API expects x_sf with trailing dim: (M, K//16, 1).
        x_sf = block_scale.unsqueeze(-1)

        # The kernel defaults swiglu_{alpha,beta,limit} to the plain-SwiGLU
        # values, so only forward the ones the model actually sets.
        swiglu_params: dict[str, float | None] = {}
        if activation == MoEActivation.SILU:
            swiglu_params = {"swiglu_limit": self.gemm1_clamp_limit}
        elif activation in (
            MoEActivation.SWIGLUOAI,
            MoEActivation.SWIGLUOAI_UNINTERLEAVE,
        ):
            swiglu_params = {
                "swiglu_alpha": self.gemm1_alpha,
                "swiglu_beta": self.gemm1_beta,
                "swiglu_limit": self.gemm1_clamp_limit,
            }
        elif activation == MoEActivation.SITU:
            # The cute_dsl kernel keys SiTU on situ_beta and requires
            # activation_type to stay a base type (ActivationType.Situ is
            # rejected by normalize_cute_dsl_moe_activation_type), so the
            # Swiglu base type is passed below and SiTU rides the betas.
            if self.situ_beta is None:
                raise ValueError(
                    "SITU activation requires moe_config.activation_situ_beta"
                )
            swiglu_params = {
                "situ_beta": self.situ_beta,
                "situ_linear_beta": self.situ_linear_beta,
            }
        swiglu_kwargs = {k: v for k, v in swiglu_params.items() if v is not None}

        per_token_kwargs = (
            {"per_token_scale": per_token_scale} if self.per_token_activation else {}
        )
        flashinfer_cute_dsl_fused_moe_nvfp4(
            x=hidden_states,
            x_sf=x_sf,
            token_selected_experts=topk_ids.to(torch.int32),
            token_final_scales=topk_weights.float(),
            w1_weight=w1,
            w1_weight_sf=self.w1_scale,
            w1_alpha=self.g1_alphas,
            fc2_input_scale=fc2_input_scale,
            w2_weight=w2,
            w2_weight_sf=self.w2_scale,
            w2_alpha=self.g2_alphas,
            num_experts=self.global_num_experts,
            top_k=self.topk,
            num_local_experts=self.local_num_experts,
            local_expert_offset=self.local_expert_offset,
            moe_output=output,
            activation_type=activation_to_flashinfer_int(
                MoEActivation.SILU if activation == MoEActivation.SITU else activation
            ),
            **swiglu_kwargs,
            **per_token_kwargs,
        )