[Step3p5] Optimize allreduce in MoE layers (#22773)

This commit is contained in:
Yuhao Yang
2026-04-16 09:33:12 +08:00
committed by GitHub
parent a4cf2ea128
commit b8794baa6d
+59 -57
View File
@@ -1,5 +1,3 @@
import logging
import os
from typing import Any, Dict, Iterable, Optional, Tuple, Union
import torch
@@ -57,7 +55,6 @@ from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, mak
Step3p5Config = None
logger = logging.getLogger(__name__)
_is_cuda = is_cuda()
@@ -69,6 +66,9 @@ class Step3p5MLP(nn.Module):
swiglu_limit: Optional[float] = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
tp_size: Optional[int] = None,
tp_rank: Optional[int] = None,
reduce_results: bool = True,
) -> None:
super().__init__()
self.hidden_size = hidden_size
@@ -79,6 +79,8 @@ class Step3p5MLP(nn.Module):
bias=False,
quant_config=quant_config,
prefix=add_prefix("gate_up_proj", prefix),
tp_size=tp_size,
tp_rank=tp_rank,
)
self.down_proj = RowParallelLinear(
intermediate_size,
@@ -86,6 +88,9 @@ class Step3p5MLP(nn.Module):
bias=False,
quant_config=quant_config,
prefix=add_prefix("down_proj", prefix),
tp_size=tp_size,
tp_rank=tp_rank,
reduce_results=reduce_results,
)
self.act_fn = SiluAndMul()
self.limit = swiglu_limit
@@ -392,6 +397,7 @@ class Step3p5Attention(nn.Module):
quant_config=quant_config,
tp_rank=attn_tp_rank,
tp_size=attn_tp_size,
reduce_results=False,
prefix=add_prefix("o_proj", prefix),
)
@@ -483,10 +489,12 @@ class Step3p5DecoderLayer(nn.Module):
rope_theta = config.rope_theta
max_position_embeddings = config.max_position_embeddings
head_dim = config.head_dim
moe_layers_list = [int(x) for x in config.moe_layers_enum.split(",")]
moe_layers_set = {int(x) for x in config.moe_layers_enum.split(",")}
self.num_attention_heads = config.num_attention_heads
self.num_key_value_heads = config.num_attention_groups
self.is_moe_layer = layer_id in moe_layers_list
self.is_moe_layer = layer_id in moe_layers_set
self.is_previous_layer_sparse = (layer_id - 1) in moe_layers_set
self.is_next_layer_sparse = (layer_id + 1) in moe_layers_set
num_hidden_layers = config.num_hidden_layers
if (
@@ -540,12 +548,16 @@ class Step3p5DecoderLayer(nn.Module):
quant_config=quant_config,
prefix=add_prefix("mlp", prefix),
)
# reduce_results=False: share_expert output stays unreduced and is
# combined with the (also unreduced) MoE output, then a single
# all-reduce covers both — saving one full-TP all-reduce per layer.
self.share_expert = Step3p5MLP(
hidden_size=self.hidden_size,
intermediate_size=config.share_expert_dim,
swiglu_limit=swiglu_limit_shared,
quant_config=quant_config,
prefix=add_prefix("share_expert", prefix),
reduce_results=False,
)
self.use_moe = True
else:
@@ -567,44 +579,19 @@ class Step3p5DecoderLayer(nn.Module):
num_layers=(
config.num_hidden_layers if layer_id < config.num_hidden_layers else 1
), # 1 is for mtp
is_layer_sparse=False,
is_previous_layer_sparse=False,
is_next_layer_sparse=False,
is_layer_sparse=self.is_moe_layer,
is_previous_layer_sparse=self.is_previous_layer_sparse,
is_next_layer_sparse=self.is_next_layer_sparse,
)
self.layer_communicator = LayerCommunicator(
layer_scatter_modes=self.layer_scatter_modes,
input_layernorm=self.input_layernorm,
post_attention_layernorm=self.post_attention_layernorm,
allow_reduce_scatter=True,
is_last_layer=(layer_id == config.num_hidden_layers - 1),
)
self.layer_id = layer_id
self.dump_intermediate = (
os.environ.get("SGLANG_DUMP_STEP3P5_INTERMEDIATE") == "1"
)
self._dump_step = 0
def _dump_tensor(
self,
name: str,
tensor: Optional[torch.Tensor],
step_id: Optional[int] = None,
) -> None:
if not self.dump_intermediate or tensor is None or not torch.is_tensor(tensor):
return
dump_dir = "/sgl-workspace/sgl"
try:
os.makedirs(dump_dir, exist_ok=True)
tp_rank = get_tensor_model_parallel_rank()
step_part = f"_step{step_id}" if step_id is not None else ""
path = os.path.join(
dump_dir,
f"step3p5_layer{self.layer_id}{step_part}_{name}_tp{tp_rank}.pt",
)
torch.save(tensor.detach().cpu(), path)
except Exception:
logger.exception(
"Failed to dump tensor %s for layer %s", name, self.layer_id
)
def forward(
self,
@@ -621,40 +608,55 @@ class Step3p5DecoderLayer(nn.Module):
forward_batch,
post_residual_addition=post_residual_addition,
)
dump_step = None
if self.dump_intermediate:
dump_step = self._dump_step
self._dump_step += 1
self._dump_tensor("attn_input", hidden_states, dump_step)
if hidden_states.shape[0] != 0:
hidden_states = self.self_attn(
positions=positions,
hidden_states=hidden_states,
forward_batch=forward_batch,
)
self._dump_tensor("attn_output", hidden_states, dump_step)
# Fully Connected
# hidden_states, residual = self.layer_communicator.prepare_mlp(
# hidden_states,
# residual,
# forward_batch,
# )
hidden_states = residual + hidden_states
residual = hidden_states
self._dump_tensor("post_attn_residual", hidden_states, dump_step)
hidden_states = self.post_attention_layernorm(hidden_states)
self._dump_tensor("mlp_input", hidden_states, dump_step)
hidden_states, residual = self.layer_communicator.prepare_mlp(
hidden_states,
residual,
forward_batch,
)
should_allreduce_fusion = (
self.layer_communicator.should_fuse_mlp_allreduce_with_next_layer(
forward_batch
)
)
use_reduce_scatter = self.layer_communicator.should_use_reduce_scatter(
forward_batch
)
if self.use_moe:
# Both share_expert and MoE return unreduced (TP-partial) outputs.
# Combine them first, then do a single all-reduce — saving one
# full-TP all-reduce per layer.
share_output = self.share_expert(hidden_states)
moe_output = self.moe(hidden_states)
moe_output = self.moe(
hidden_states,
forward_batch,
should_allreduce_fusion=True,
use_reduce_scatter=use_reduce_scatter,
)
hidden_states = moe_output + share_output
if not should_allreduce_fusion and not use_reduce_scatter:
hidden_states = tensor_model_parallel_all_reduce(hidden_states)
else:
hidden_states = self.mlp(hidden_states)
self._dump_tensor("mlp_output", hidden_states, dump_step)
hidden_states, residual = self.layer_communicator.postprocess_layer(
hidden_states, residual, forward_batch
)
self._dump_tensor("layer_output", hidden_states, dump_step)
# Dense MLP uses reduce_results=True, so the output is already
# all-reduced. Do NOT set the fusion flag — otherwise the next
# layer would all-reduce again, multiplying values by world_size.
should_allreduce_fusion = False
if should_allreduce_fusion:
hidden_states._sglang_needs_allreduce_fusion = True
else:
hidden_states, residual = self.layer_communicator.postprocess_layer(
hidden_states, residual, forward_batch
)
return hidden_states, residual