Fix DP attention + EP mode of Nemotron (#28102)
Co-authored-by: Brayden Zhong <brayden@radixark.ai>
This commit is contained in:
co-authored by
Brayden Zhong
parent
6c3e429ba1
commit
1a19f66acb
@@ -55,7 +55,10 @@ from sglang.srt.layers.logits_processor import LogitsProcessor
|
|||||||
from sglang.srt.layers.moe.ep_moe.layer import get_moe_impl_class
|
from sglang.srt.layers.moe.ep_moe.layer import get_moe_impl_class
|
||||||
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
|
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
|
||||||
from sglang.srt.layers.moe.topk import TopK
|
from sglang.srt.layers.moe.topk import TopK
|
||||||
from sglang.srt.layers.moe.utils import RoutingMethodType
|
from sglang.srt.layers.moe.utils import (
|
||||||
|
RoutingMethodType,
|
||||||
|
should_skip_post_experts_all_reduce,
|
||||||
|
)
|
||||||
from sglang.srt.layers.quantization import QuantizationConfig
|
from sglang.srt.layers.quantization import QuantizationConfig
|
||||||
from sglang.srt.layers.radix_attention import RadixAttention
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
from sglang.srt.layers.utils import PPMissingLayer, get_layer_id
|
from sglang.srt.layers.utils import PPMissingLayer, get_layer_id
|
||||||
@@ -129,10 +132,10 @@ class NemotronHMLP(nn.Module):
|
|||||||
)
|
)
|
||||||
self.act_fn = ReLU2()
|
self.act_fn = ReLU2()
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor):
|
def forward(self, x: torch.Tensor, use_reduce_scatter: bool = False):
|
||||||
x, _ = self.up_proj(x)
|
x, _ = self.up_proj(x)
|
||||||
x = self.act_fn(x)
|
x = self.act_fn(x)
|
||||||
x, _ = self.down_proj(x)
|
x, _ = self.down_proj(x, skip_all_reduce=use_reduce_scatter)
|
||||||
return x
|
return x
|
||||||
|
|
||||||
|
|
||||||
@@ -296,7 +299,9 @@ class NemotronHMoE(nn.Module):
|
|||||||
|
|
||||||
return final_hidden_states, shared_output
|
return final_hidden_states, shared_output
|
||||||
|
|
||||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
def forward(
|
||||||
|
self, hidden_states: torch.Tensor, use_reduce_scatter: bool = False
|
||||||
|
) -> torch.Tensor:
|
||||||
num_tokens, hidden_dim = hidden_states.shape
|
num_tokens, hidden_dim = hidden_states.shape
|
||||||
# routed_scaling_factor is fused into the experts call (applied by the
|
# routed_scaling_factor is fused into the experts call (applied by the
|
||||||
# MoE runner / topk), so final_hidden_states is already scaled.
|
# MoE runner / topk), so final_hidden_states is already scaled.
|
||||||
@@ -308,7 +313,10 @@ class NemotronHMoE(nn.Module):
|
|||||||
if shared_output is not None:
|
if shared_output is not None:
|
||||||
final_hidden_states += shared_output
|
final_hidden_states += shared_output
|
||||||
|
|
||||||
if self.tp_size > 1:
|
if self.tp_size > 1 and not should_skip_post_experts_all_reduce(
|
||||||
|
is_tp_path=True,
|
||||||
|
use_reduce_scatter=use_reduce_scatter,
|
||||||
|
):
|
||||||
final_hidden_states = tensor_model_parallel_all_reduce(final_hidden_states)
|
final_hidden_states = tensor_model_parallel_all_reduce(final_hidden_states)
|
||||||
|
|
||||||
return final_hidden_states.view(num_tokens, hidden_dim)
|
return final_hidden_states.view(num_tokens, hidden_dim)
|
||||||
@@ -328,7 +336,12 @@ class NemotronHMLPLikeDecoderLayer(nn.Module):
|
|||||||
hidden_states, residual = self.layer_communicator.prepare_mlp(
|
hidden_states, residual = self.layer_communicator.prepare_mlp(
|
||||||
hidden_states, residual, forward_batch
|
hidden_states, residual, forward_batch
|
||||||
)
|
)
|
||||||
hidden_states = self.mixer.forward(hidden_states)
|
use_reduce_scatter = self.layer_communicator.should_use_reduce_scatter(
|
||||||
|
forward_batch
|
||||||
|
)
|
||||||
|
hidden_states = self.mixer.forward(
|
||||||
|
hidden_states, use_reduce_scatter=use_reduce_scatter
|
||||||
|
)
|
||||||
hidden_states, residual = self.layer_communicator.postprocess_layer(
|
hidden_states, residual = self.layer_communicator.postprocess_layer(
|
||||||
hidden_states, residual, forward_batch
|
hidden_states, residual, forward_batch
|
||||||
)
|
)
|
||||||
@@ -375,7 +388,9 @@ class NemotronHMLPDecoderLayer(NemotronHMLPLikeDecoderLayer):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
||||||
self.layer_communicator = make_layer_communicator(self.norm, for_attn=False)
|
self.layer_communicator = make_layer_communicator(
|
||||||
|
self.norm, for_attn=False, allow_reduce_scatter=True
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class NemotronHMoEDecoderLayer(NemotronHMLPLikeDecoderLayer):
|
class NemotronHMoEDecoderLayer(NemotronHMLPLikeDecoderLayer):
|
||||||
@@ -398,7 +413,9 @@ class NemotronHMoEDecoderLayer(NemotronHMLPLikeDecoderLayer):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
||||||
self.layer_communicator = make_layer_communicator(self.norm, for_attn=False)
|
self.layer_communicator = make_layer_communicator(
|
||||||
|
self.norm, for_attn=False, allow_reduce_scatter=True
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class NemotronHAttnLikeDecoderLayer(nn.Module):
|
class NemotronHAttnLikeDecoderLayer(nn.Module):
|
||||||
|
|||||||
@@ -57,11 +57,12 @@ def _build_layer_scatter_modes() -> LayerScatterModes:
|
|||||||
|
|
||||||
|
|
||||||
def make_layer_communicator(
|
def make_layer_communicator(
|
||||||
layer_norm: RMSNorm, *, for_attn: bool
|
layer_norm: RMSNorm, *, for_attn: bool, allow_reduce_scatter: bool = False
|
||||||
) -> LayerCommunicator:
|
) -> LayerCommunicator:
|
||||||
return LayerCommunicator(
|
return LayerCommunicator(
|
||||||
layer_scatter_modes=_build_layer_scatter_modes(),
|
layer_scatter_modes=_build_layer_scatter_modes(),
|
||||||
input_layernorm=layer_norm if for_attn else nn.Identity(),
|
input_layernorm=layer_norm if for_attn else nn.Identity(),
|
||||||
post_attention_layernorm=nn.Identity() if for_attn else layer_norm,
|
post_attention_layernorm=nn.Identity() if for_attn else layer_norm,
|
||||||
force_layernorm_before_dp_gather=True,
|
force_layernorm_before_dp_gather=True,
|
||||||
|
allow_reduce_scatter=allow_reduce_scatter,
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user