Skip to content

vllm.model_executor.layers.quantization.online.moe_base

Classes:

OnlineMoEMethodBase

Bases: FusedMoEMethodBase

Base for MoE methods that load full-precision weights on meta device and quantize them after loading via the QeRL layerwise processing system.

Source code in vllm/model_executor/layers/quantization/online/moe_base.py
class OnlineMoEMethodBase(FusedMoEMethodBase):
    """Base for MoE methods that load full-precision weights on meta device
    and quantize them after loading via the QeRL layerwise processing system.
    """

    uses_meta_device: bool = True

    def create_weights(
        self,
        layer: torch.nn.Module,
        num_experts: int,
        hidden_size: int,
        intermediate_size_per_partition: int,
        params_dtype: torch.dtype,
        **extra_weight_attrs,
    ):
        layer.num_experts = num_experts
        layer.orig_dtype = params_dtype
        layer.weight_block_size = None

        # Fused gate_up_proj (column parallel) — full precision on meta device
        w13_weight = torch.nn.Parameter(
            torch.empty(
                num_experts,
                self.moe.w13_num_shards * intermediate_size_per_partition,
                hidden_size,
                device="meta",
                dtype=params_dtype,
            ),
            requires_grad=False,
        )
        layer.register_parameter("w13_weight", w13_weight)
        set_weight_attrs(w13_weight, extra_weight_attrs)

        # down_proj (row parallel) — full precision on meta device
        w2_weight = torch.nn.Parameter(
            torch.empty(
                num_experts,
                hidden_size,
                intermediate_size_per_partition,
                device="meta",
                dtype=params_dtype,
            ),
            requires_grad=False,
        )
        layer.register_parameter("w2_weight", w2_weight)
        set_weight_attrs(w2_weight, extra_weight_attrs)

        # BIASES (for models like GPT-OSS that have biased MoE)
        if self.moe.has_bias:
            w13_bias = torch.nn.Parameter(
                torch.zeros(
                    num_experts,
                    self.moe.w13_num_shards * intermediate_size_per_partition,
                    device="meta",
                    dtype=layer.orig_dtype,
                ),
                requires_grad=False,
            )
            layer.register_parameter("w13_bias", w13_bias)
            set_weight_attrs(w13_bias, extra_weight_attrs)

            w2_bias = torch.nn.Parameter(
                torch.zeros(
                    num_experts,
                    hidden_size,
                    device="meta",
                    dtype=layer.orig_dtype,
                ),
                requires_grad=False,
            )
            layer.register_parameter("w2_bias", w2_bias)
            set_weight_attrs(w2_bias, extra_weight_attrs)

        layer.w13_input_scale = None
        layer.w2_input_scale = None

        initialize_online_processing(layer)

    def _zero_padding(self, layer: torch.nn.Module) -> None:
        hidden_size = layer.moe_config.hidden_dim_unpadded
        intermediate_size = layer.moe_config.intermediate_size_per_partition_unpadded

        w13_shard = layer.w13_weight.shape[1] // self.moe.w13_num_shards
        if w13_shard > intermediate_size:
            for shard in range(self.moe.w13_num_shards):
                start = shard * w13_shard + intermediate_size
                layer.w13_weight[:, start : (shard + 1) * w13_shard, :] = 0
        if layer.w13_weight.shape[2] > hidden_size:
            layer.w13_weight[:, :, hidden_size:] = 0

        if layer.w2_weight.shape[1] > hidden_size:
            layer.w2_weight[:, hidden_size:, :] = 0
        if layer.w2_weight.shape[2] > intermediate_size:
            layer.w2_weight[:, :, intermediate_size:] = 0

        if getattr(layer, "w13_bias", None) is not None:
            w13_bias_shard = layer.w13_bias.shape[1] // self.moe.w13_num_shards
            if w13_bias_shard > intermediate_size:
                for shard in range(self.moe.w13_num_shards):
                    start = shard * w13_bias_shard + intermediate_size
                    layer.w13_bias[:, start : (shard + 1) * w13_bias_shard] = 0

        if (
            getattr(layer, "w2_bias", None) is not None
            and layer.w2_bias.shape[1] > hidden_size
        ):
            layer.w2_bias[:, hidden_size:] = 0

    @abstractmethod
    def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
        pass

    @property
    def supports_eplb(self) -> bool:
        return True

    def apply_monolithic(
        self,
        layer: RoutedExperts,
        x: torch.Tensor,
        router_logits: torch.Tensor,
        input_ids: torch.Tensor | None = None,
    ) -> torch.Tensor:
        assert self.is_monolithic
        assert self.moe_kernel is not None
        return self.moe_kernel.apply_monolithic(
            x,
            layer.w13_weight,
            layer.w2_weight,
            router_logits,
            activation=layer.activation,
            global_num_experts=layer.global_num_experts,
            expert_map=layer.expert_map,
            apply_router_weight_on_input=layer.apply_router_weight_on_input,
            num_expert_group=layer.num_expert_group,
            topk_group=layer.topk_group,
            e_score_correction_bias=layer.e_score_correction_bias,
            routed_scaling_factor=layer.routed_scaling_factor,
        )

    def apply(
        self,
        layer: RoutedExperts,
        x: torch.Tensor,
        topk_weights: torch.Tensor,
        topk_ids: torch.Tensor,
        shared_experts: SharedExperts | None,
        shared_experts_input: torch.Tensor | None,
    ) -> torch.Tensor:
        assert not self.is_monolithic
        assert self.moe_kernel is not None
        return self.moe_kernel.apply(
            x,
            layer.w13_weight,
            layer.w2_weight,
            topk_weights,
            topk_ids,
            activation=layer.activation,
            global_num_experts=layer.global_num_experts,
            expert_map=layer.expert_map,
            apply_router_weight_on_input=layer.apply_router_weight_on_input,
            shared_experts=shared_experts,
            shared_experts_input=shared_experts_input,
        )