Fix DP attention + EP mode of Nemotron (#28102)

Co-authored-by: Brayden Zhong <brayden@radixark.ai>
This commit is contained in:
Brayden Zhong
2026-06-12 18:39:04 -07:00
committed by GitHub
co-authored by Brayden Zhong
parent 6c3e429ba1
commit 1a19f66acb
2 changed files with 27 additions and 9 deletions
+25 -8
View File
@@ -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.fused_moe_triton.layer import FusedMoE
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.radix_attention import RadixAttention
from sglang.srt.layers.utils import PPMissingLayer, get_layer_id
@@ -129,10 +132,10 @@ class NemotronHMLP(nn.Module):
)
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.act_fn(x)
x, _ = self.down_proj(x)
x, _ = self.down_proj(x, skip_all_reduce=use_reduce_scatter)
return x
@@ -296,7 +299,9 @@ class NemotronHMoE(nn.Module):
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
# routed_scaling_factor is fused into the experts call (applied by the
# MoE runner / topk), so final_hidden_states is already scaled.
@@ -308,7 +313,10 @@ class NemotronHMoE(nn.Module):
if shared_output is not None:
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)
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, 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, forward_batch
)
@@ -375,7 +388,9 @@ class NemotronHMLPDecoderLayer(NemotronHMLPLikeDecoderLayer):
)
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):
@@ -398,7 +413,9 @@ class NemotronHMoEDecoderLayer(NemotronHMLPLikeDecoderLayer):
)
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):
+2 -1
View File
@@ -57,11 +57,12 @@ def _build_layer_scatter_modes() -> LayerScatterModes:
def make_layer_communicator(
layer_norm: RMSNorm, *, for_attn: bool
layer_norm: RMSNorm, *, for_attn: bool, allow_reduce_scatter: bool = False
) -> LayerCommunicator:
return LayerCommunicator(
layer_scatter_modes=_build_layer_scatter_modes(),
input_layernorm=layer_norm if for_attn else nn.Identity(),
post_attention_layernorm=nn.Identity() if for_attn else layer_norm,
force_layernorm_before_dp_gather=True,
allow_reduce_scatter=allow_reduce_scatter,
)