[Step3p5] Optimize allreduce in MoE layers (#22773)
This commit is contained in:
@@ -1,5 +1,3 @@
|
|||||||
import logging
|
|
||||||
import os
|
|
||||||
from typing import Any, Dict, Iterable, Optional, Tuple, Union
|
from typing import Any, Dict, Iterable, Optional, Tuple, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -57,7 +55,6 @@ from sglang.srt.utils import add_prefix, is_cuda, is_non_idle_and_non_empty, mak
|
|||||||
|
|
||||||
Step3p5Config = None
|
Step3p5Config = None
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
|
|
||||||
|
|
||||||
@@ -69,6 +66,9 @@ class Step3p5MLP(nn.Module):
|
|||||||
swiglu_limit: Optional[float] = None,
|
swiglu_limit: Optional[float] = None,
|
||||||
quant_config: Optional[QuantizationConfig] = None,
|
quant_config: Optional[QuantizationConfig] = None,
|
||||||
prefix: str = "",
|
prefix: str = "",
|
||||||
|
tp_size: Optional[int] = None,
|
||||||
|
tp_rank: Optional[int] = None,
|
||||||
|
reduce_results: bool = True,
|
||||||
) -> None:
|
) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.hidden_size = hidden_size
|
self.hidden_size = hidden_size
|
||||||
@@ -79,6 +79,8 @@ class Step3p5MLP(nn.Module):
|
|||||||
bias=False,
|
bias=False,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("gate_up_proj", prefix),
|
prefix=add_prefix("gate_up_proj", prefix),
|
||||||
|
tp_size=tp_size,
|
||||||
|
tp_rank=tp_rank,
|
||||||
)
|
)
|
||||||
self.down_proj = RowParallelLinear(
|
self.down_proj = RowParallelLinear(
|
||||||
intermediate_size,
|
intermediate_size,
|
||||||
@@ -86,6 +88,9 @@ class Step3p5MLP(nn.Module):
|
|||||||
bias=False,
|
bias=False,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("down_proj", prefix),
|
prefix=add_prefix("down_proj", prefix),
|
||||||
|
tp_size=tp_size,
|
||||||
|
tp_rank=tp_rank,
|
||||||
|
reduce_results=reduce_results,
|
||||||
)
|
)
|
||||||
self.act_fn = SiluAndMul()
|
self.act_fn = SiluAndMul()
|
||||||
self.limit = swiglu_limit
|
self.limit = swiglu_limit
|
||||||
@@ -392,6 +397,7 @@ class Step3p5Attention(nn.Module):
|
|||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
tp_rank=attn_tp_rank,
|
tp_rank=attn_tp_rank,
|
||||||
tp_size=attn_tp_size,
|
tp_size=attn_tp_size,
|
||||||
|
reduce_results=False,
|
||||||
prefix=add_prefix("o_proj", prefix),
|
prefix=add_prefix("o_proj", prefix),
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -483,10 +489,12 @@ class Step3p5DecoderLayer(nn.Module):
|
|||||||
rope_theta = config.rope_theta
|
rope_theta = config.rope_theta
|
||||||
max_position_embeddings = config.max_position_embeddings
|
max_position_embeddings = config.max_position_embeddings
|
||||||
head_dim = config.head_dim
|
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_attention_heads = config.num_attention_heads
|
||||||
self.num_key_value_heads = config.num_attention_groups
|
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
|
num_hidden_layers = config.num_hidden_layers
|
||||||
|
|
||||||
if (
|
if (
|
||||||
@@ -540,12 +548,16 @@ class Step3p5DecoderLayer(nn.Module):
|
|||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("mlp", prefix),
|
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(
|
self.share_expert = Step3p5MLP(
|
||||||
hidden_size=self.hidden_size,
|
hidden_size=self.hidden_size,
|
||||||
intermediate_size=config.share_expert_dim,
|
intermediate_size=config.share_expert_dim,
|
||||||
swiglu_limit=swiglu_limit_shared,
|
swiglu_limit=swiglu_limit_shared,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("share_expert", prefix),
|
prefix=add_prefix("share_expert", prefix),
|
||||||
|
reduce_results=False,
|
||||||
)
|
)
|
||||||
self.use_moe = True
|
self.use_moe = True
|
||||||
else:
|
else:
|
||||||
@@ -567,44 +579,19 @@ class Step3p5DecoderLayer(nn.Module):
|
|||||||
num_layers=(
|
num_layers=(
|
||||||
config.num_hidden_layers if layer_id < config.num_hidden_layers else 1
|
config.num_hidden_layers if layer_id < config.num_hidden_layers else 1
|
||||||
), # 1 is for mtp
|
), # 1 is for mtp
|
||||||
is_layer_sparse=False,
|
is_layer_sparse=self.is_moe_layer,
|
||||||
is_previous_layer_sparse=False,
|
is_previous_layer_sparse=self.is_previous_layer_sparse,
|
||||||
is_next_layer_sparse=False,
|
is_next_layer_sparse=self.is_next_layer_sparse,
|
||||||
)
|
)
|
||||||
self.layer_communicator = LayerCommunicator(
|
self.layer_communicator = LayerCommunicator(
|
||||||
layer_scatter_modes=self.layer_scatter_modes,
|
layer_scatter_modes=self.layer_scatter_modes,
|
||||||
input_layernorm=self.input_layernorm,
|
input_layernorm=self.input_layernorm,
|
||||||
post_attention_layernorm=self.post_attention_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.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(
|
def forward(
|
||||||
self,
|
self,
|
||||||
@@ -621,40 +608,55 @@ class Step3p5DecoderLayer(nn.Module):
|
|||||||
forward_batch,
|
forward_batch,
|
||||||
post_residual_addition=post_residual_addition,
|
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:
|
if hidden_states.shape[0] != 0:
|
||||||
hidden_states = self.self_attn(
|
hidden_states = self.self_attn(
|
||||||
positions=positions,
|
positions=positions,
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
forward_batch=forward_batch,
|
forward_batch=forward_batch,
|
||||||
)
|
)
|
||||||
self._dump_tensor("attn_output", hidden_states, dump_step)
|
|
||||||
# Fully Connected
|
# Fully Connected
|
||||||
# hidden_states, residual = self.layer_communicator.prepare_mlp(
|
hidden_states, residual = self.layer_communicator.prepare_mlp(
|
||||||
# hidden_states,
|
hidden_states,
|
||||||
# residual,
|
residual,
|
||||||
# forward_batch,
|
forward_batch,
|
||||||
# )
|
)
|
||||||
hidden_states = residual + hidden_states
|
|
||||||
residual = hidden_states
|
should_allreduce_fusion = (
|
||||||
self._dump_tensor("post_attn_residual", hidden_states, dump_step)
|
self.layer_communicator.should_fuse_mlp_allreduce_with_next_layer(
|
||||||
hidden_states = self.post_attention_layernorm(hidden_states)
|
forward_batch
|
||||||
self._dump_tensor("mlp_input", hidden_states, dump_step)
|
)
|
||||||
|
)
|
||||||
|
use_reduce_scatter = self.layer_communicator.should_use_reduce_scatter(
|
||||||
|
forward_batch
|
||||||
|
)
|
||||||
|
|
||||||
if self.use_moe:
|
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)
|
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
|
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:
|
else:
|
||||||
hidden_states = self.mlp(hidden_states)
|
hidden_states = self.mlp(hidden_states)
|
||||||
self._dump_tensor("mlp_output", hidden_states, dump_step)
|
# Dense MLP uses reduce_results=True, so the output is already
|
||||||
hidden_states, residual = self.layer_communicator.postprocess_layer(
|
# all-reduced. Do NOT set the fusion flag — otherwise the next
|
||||||
hidden_states, residual, forward_batch
|
# layer would all-reduce again, multiplying values by world_size.
|
||||||
)
|
should_allreduce_fusion = False
|
||||||
self._dump_tensor("layer_output", hidden_states, dump_step)
|
|
||||||
|
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
|
return hidden_states, residual
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user