@CustomOp.register("unquantized_fused_moe")
class UnquantizedFusedMoEMethod(FusedMoEMethodBase, CustomOp):
"""MoE method without quantization."""
# --8<-- [end:unquantized_fused_moe]
def __init__(self, moe: FusedMoEConfig):
super().__init__(moe)
self.unquantized_backend, self.experts_cls = select_unquantized_moe_backend(
moe_config=self.moe,
)
@property
def supports_eplb(self) -> bool:
return True
def create_weights(
self,
layer: "RoutedExperts",
num_experts: int,
hidden_size: int,
intermediate_size_per_partition: int,
params_dtype: torch.dtype,
**extra_weight_attrs,
):
if self.moe.is_act_and_mul:
w13_up_dim = 2 * intermediate_size_per_partition
else:
w13_up_dim = intermediate_size_per_partition
# Fused gate_up_proj (column parallel)
w13_weight = torch.nn.Parameter(
torch.empty(
num_experts,
w13_up_dim,
hidden_size,
dtype=params_dtype,
),
requires_grad=False,
)
layer.register_parameter("w13_weight", w13_weight)
set_weight_attrs(w13_weight, extra_weight_attrs)
if self.moe.has_bias:
w13_bias = torch.nn.Parameter(
torch.zeros(num_experts, w13_up_dim, dtype=params_dtype),
requires_grad=False,
)
layer.register_parameter("w13_bias", w13_bias)
set_weight_attrs(w13_bias, extra_weight_attrs)
# down_proj (row parallel)
w2_weight = torch.nn.Parameter(
torch.empty(
num_experts,
hidden_size,
intermediate_size_per_partition,
dtype=params_dtype,
),
requires_grad=False,
)
layer.register_parameter("w2_weight", w2_weight)
set_weight_attrs(w2_weight, extra_weight_attrs)
if self.moe.has_bias:
w2_bias = torch.nn.Parameter(
torch.zeros(num_experts, hidden_size, dtype=params_dtype),
requires_grad=False,
)
layer.register_parameter("w2_bias", w2_bias)
set_weight_attrs(w2_bias, extra_weight_attrs)
def _maybe_pad_weight(self, weight: torch.Tensor) -> torch.Tensor:
# Pad the weight tensor. This is an optimization on ROCm platform, which
# can benefit from tensors located far enough from one another in memory.
# Skip padding when EPLB is enabled because EPLB requires contiguous
# weights for the view/rearrangement operations.
if (
envs.VLLM_ROCM_MOE_PADDING
and current_platform.is_rocm()
and not self.moe.moe_parallel_config.enable_eplb
and weight.stride(-1) == 1
and (weight.stride(-2) * weight.element_size()) % 512 == 0
):
num_pad = 256 // weight.element_size()
weight = F.pad(weight, (0, num_pad), "constant", 0)[..., :-num_pad]
torch.accelerator.empty_cache()
return weight
def _setup_kernel(
self,
layer: "RoutedExperts",
w13: torch.Tensor,
w2: torch.Tensor,
) -> None:
# Shuffle weights to runtime format.
w13_new, w2_new = convert_to_unquantized_kernel_format(
self.unquantized_backend,
moe_config=layer.moe_config,
w13_weight=w13,
w2_weight=w2,
)
# `moe_kernel` is initialized to None in FusedMoEMethodBase.__init__;
# On the first call we replace the parameter normally. On subsequent
# calls (e.g. RL weight updates that re-trigger
# process_weights_after_loading) the moe kernel has already been set
# up and CUDA graphs may have captured the parameter addresses, so
# we copy the shuffled data into the existing storage instead of
# re-registering a new Parameter.
is_weight_update = self.moe_kernel is not None # type: ignore[has-type]
replace_parameter(layer, "w13_weight", w13_new, prefer_copy=is_weight_update)
replace_parameter(layer, "w2_weight", w2_new, prefer_copy=is_weight_update)
if not is_weight_update:
# Setup moe kernel only on the first call. For the unquantized
# method, moe_quant_config carries no quantized scales -- only
# optional w{13,2}_bias references and SwiGLU gate params. Since
# weight updates mutate those bias tensors in place, the kernel
# does not need to be re-built.
self.moe_quant_config = self.get_fused_moe_quant_config(layer)
assert self.moe_quant_config is not None
assert self.experts_cls is not None
self.moe_kernel = make_unquantized_moe_kernel(
quant_config=self.moe_quant_config,
moe_config=self.moe,
backend=self.unquantized_backend,
experts_cls=self.experts_cls,
routing_tables=layer._expert_routing_tables(),
)
if self.unquantized_backend == UnquantizedMoeBackend.CPU:
# The CPU experts need the layer itself for the setup that
# convert_to_unquantized_kernel_format cannot express, since
# it only sees the two weight tensors: padding and prepacking
# into the grouped-gemm layout (bias included), and capturing
# the router config that monolithic apply() cannot carry.
self.moe_kernel.fused_experts.process_weights_after_loading(layer)
def process_weights_after_loading(self, layer: "RoutedExperts") -> None:
super().process_weights_after_loading(layer)
# Padding the weight for better performance on ROCm.
# _maybe_pad_weight is idempotent: on the first call it allocates a
# padded storage and returns a strided view; on subsequent calls
# (weight updates) the stride condition no longer matches so it
# returns the input unchanged. The reassignment to .data is therefore
# a no-op on updates and preserves the storage address (data_ptr)
# used by captured CUDA graphs.
layer.w13_weight.data = self._maybe_pad_weight(layer.w13_weight.data)
layer.w2_weight.data = self._maybe_pad_weight(layer.w2_weight.data)
if self.unquantized_backend in [
UnquantizedMoeBackend.TPU,
UnquantizedMoeBackend.OOT,
]:
# OOT handles internally.
return
elif self.unquantized_backend == UnquantizedMoeBackend.XPU:
w13 = layer.w13_weight
w2 = layer.w2_weight
w13.data = w13.transpose(-1, -2).contiguous()
w2.data = w2.transpose(-1, -2).contiguous()
self._setup_kernel(
layer=layer,
w13=w13,
w2=w2,
)
else:
self._setup_kernel(
layer=layer,
w13=layer.w13_weight,
w2=layer.w2_weight,
)
def get_fused_moe_quant_config(self, layer: torch.nn.Module) -> FusedMoEQuantConfig:
# SwiGLU/swigluoai gate params live on the layer; plumb them into the
# quant config so the fused activation (e.g. swigluoai_uninterleave on
# MiniMax-M3) receives gemm1_clamp_limit/alpha/beta.
gemm1_alpha = getattr(layer, "swiglu_alpha", None)
gemm1_beta = getattr(layer, "swiglu_beta", None)
gemm1_clamp_limit = getattr(layer, "swiglu_limit", None)
if self.moe.has_bias:
return biased_moe_quant_config(
layer.w13_bias,
layer.w2_bias,
gemm1_alpha=gemm1_alpha,
gemm1_beta=gemm1_beta,
gemm1_clamp_limit=gemm1_clamp_limit,
)
return FusedMoEQuantConfig.make(
gemm1_alpha=gemm1_alpha,
gemm1_beta=gemm1_beta,
gemm1_clamp_limit=gemm1_clamp_limit,
)
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:
return self.forward(
layer=layer,
x=x,
topk_weights=topk_weights,
topk_ids=topk_ids,
shared_experts=shared_experts,
shared_experts_input=shared_experts_input,
)
def forward_native(
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 self.moe_kernel is not None
return self.moe_kernel.apply(
hidden_states=x,
w1=layer.w13_weight,
w2=layer.w2_weight,
topk_weights=topk_weights,
topk_ids=topk_ids,
activation=layer.activation,
apply_router_weight_on_input=layer.apply_router_weight_on_input,
global_num_experts=layer.global_num_experts,
expert_map=layer.expert_map,
shared_experts=shared_experts,
shared_experts_input=shared_experts_input,
)
def forward_cuda(
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:
return self.forward_native(
layer,
x,
topk_weights,
topk_ids,
shared_experts,
shared_experts_input,
)
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,
)