Use Flashinfer allreduce fusion for MNNVL allreduce for Nemotron (#28346)

Co-authored-by: Brayden Zhong <brayden@radixark.ai>
This commit is contained in:
Brayden Zhong
2026-06-18 17:30:12 -07:00
committed by GitHub
co-authored by Brayden Zhong
parent 62ab09a478
commit ea407df4b0
4 changed files with 111 additions and 47 deletions
+76 -46
View File
@@ -18,7 +18,6 @@
"""Inference-only NemotronH model."""
from collections.abc import Iterable
from typing import Optional, Union
import torch
from torch import nn
@@ -82,6 +81,7 @@ from sglang.srt.model_loader.weight_utils import (
)
from sglang.srt.models.nemotron_h_utils import (
get_real_num_tokens,
input_norm_maybe_fuse_allreduce,
is_attn_layer,
make_layer_communicator,
pad_to_original_num_tokens,
@@ -106,7 +106,7 @@ class NemotronHMLP(nn.Module):
self,
config: NemotronHConfig,
intermediate_size: int,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
bias: bool = False,
reduce_results: bool = True,
prefix: str = "",
@@ -130,10 +130,17 @@ class NemotronHMLP(nn.Module):
)
self.act_fn = ReLU2()
def forward(self, x: torch.Tensor, use_reduce_scatter: bool = False):
def forward(
self,
x: torch.Tensor,
should_allreduce_fusion: bool = False,
use_reduce_scatter: bool = False,
):
x, _ = self.up_proj(x)
x = self.act_fn(x)
x, _ = self.down_proj(x, skip_all_reduce=use_reduce_scatter)
x, _ = self.down_proj(
x, skip_all_reduce=should_allreduce_fusion or use_reduce_scatter
)
return x
@@ -152,7 +159,7 @@ class NemotronHMoE(nn.Module):
self,
config: NemotronHConfig,
layer_idx: int,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
) -> None:
super().__init__()
@@ -298,7 +305,10 @@ class NemotronHMoE(nn.Module):
return final_hidden_states, shared_output
def forward(
self, hidden_states: torch.Tensor, use_reduce_scatter: bool = False
self,
hidden_states: torch.Tensor,
should_allreduce_fusion: bool = False,
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
@@ -314,6 +324,7 @@ class NemotronHMoE(nn.Module):
if self.tp_size > 1 and not should_skip_post_experts_all_reduce(
is_tp_path=True,
use_reduce_scatter=use_reduce_scatter,
should_allreduce_fusion=should_allreduce_fusion,
):
final_hidden_states = tensor_model_parallel_all_reduce(final_hidden_states)
@@ -327,7 +338,7 @@ class NemotronHMLPLikeDecoderLayer(nn.Module):
self,
*,
hidden_states: torch.Tensor,
residual: Optional[torch.Tensor],
residual: torch.Tensor | None,
forward_batch: ForwardBatch,
) -> tuple[torch.Tensor, torch.Tensor]:
if is_dp_attention_enabled():
@@ -337,21 +348,38 @@ class NemotronHMLPLikeDecoderLayer(nn.Module):
use_reduce_scatter = self.layer_communicator.should_use_reduce_scatter(
forward_batch
)
should_allreduce_fusion = (
self.layer_communicator.should_fuse_mlp_allreduce_with_next_layer(
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
hidden_states,
should_allreduce_fusion=should_allreduce_fusion,
use_reduce_scatter=use_reduce_scatter,
)
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
if residual is None:
residual = hidden_states
hidden_states = self.norm(hidden_states)
else:
hidden_states, residual = self.norm(hidden_states, residual)
hidden_states, residual = input_norm_maybe_fuse_allreduce(
self.norm, hidden_states, residual
)
hidden_states = self.mixer.forward(hidden_states)
should_allreduce_fusion = (
self.layer_communicator.should_fuse_mlp_allreduce_with_next_layer(
forward_batch
)
)
hidden_states = self.mixer.forward(
hidden_states, should_allreduce_fusion=should_allreduce_fusion
)
if should_allreduce_fusion:
hidden_states._sglang_needs_allreduce_fusion = True
return hidden_states, residual
@@ -360,7 +388,7 @@ class NemotronHMLPDecoderLayer(NemotronHMLPLikeDecoderLayer):
self,
config: NemotronHConfig,
layer_idx: int,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
) -> None:
super().__init__()
@@ -387,7 +415,10 @@ 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, allow_reduce_scatter=True
self.norm,
for_attn=False,
allow_reduce_scatter=True,
is_last_layer=layer_idx == len(config.hybrid_override_pattern) - 1,
)
@@ -396,7 +427,7 @@ class NemotronHMoEDecoderLayer(NemotronHMLPLikeDecoderLayer):
self,
config: NemotronHConfig,
layer_idx: int,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
) -> None:
super().__init__()
@@ -412,7 +443,10 @@ 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, allow_reduce_scatter=True
self.norm,
for_attn=False,
allow_reduce_scatter=True,
is_last_layer=layer_idx == len(config.hybrid_override_pattern) - 1,
)
@@ -427,9 +461,9 @@ class NemotronHAttnLikeDecoderLayer(nn.Module):
def _dp_attn_input(
self,
hidden_states: torch.Tensor,
residual: Optional[torch.Tensor],
residual: torch.Tensor | None,
forward_batch: ForwardBatch,
) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
) -> tuple[torch.Tensor, torch.Tensor | None]:
if self.prev_layer_is_attn and residual is not None:
hidden_states = attn_tp_all_reduce(hidden_states)
return self.layer_communicator.prepare_attn(
@@ -442,7 +476,7 @@ class NemotronHMambaDecoderLayer(NemotronHAttnLikeDecoderLayer):
self,
config: NemotronHConfig,
layer_idx: int,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
) -> None:
super().__init__()
@@ -490,7 +524,7 @@ class NemotronHMambaDecoderLayer(NemotronHAttnLikeDecoderLayer):
self,
*,
hidden_states: torch.Tensor,
residual: Optional[torch.Tensor],
residual: torch.Tensor | None,
forward_batch: ForwardBatch,
) -> tuple[torch.Tensor, torch.Tensor]:
if is_dp_attention_enabled():
@@ -506,11 +540,9 @@ class NemotronHMambaDecoderLayer(NemotronHAttnLikeDecoderLayer):
output = self._forward_mamba(hidden_states, forward_batch)
return output, residual
if residual is None:
residual = hidden_states
hidden_states = self.norm(hidden_states)
else:
hidden_states, residual = self.norm(hidden_states, residual)
hidden_states, residual = input_norm_maybe_fuse_allreduce(
self.norm, hidden_states, residual
)
if is_in_breakable_cuda_graph():
output = torch.empty_like(hidden_states)
@@ -531,7 +563,7 @@ class NemotronHAttention(nn.Module):
self,
config: NemotronHConfig,
layer_idx: int,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
) -> None:
super().__init__()
@@ -636,7 +668,7 @@ class NemotronHAttentionDecoderLayer(NemotronHAttnLikeDecoderLayer):
self,
config: NemotronHConfig,
layer_idx: int,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
) -> None:
super().__init__()
@@ -657,7 +689,7 @@ class NemotronHAttentionDecoderLayer(NemotronHAttnLikeDecoderLayer):
self,
*,
hidden_states: torch.Tensor,
residual: Optional[torch.Tensor],
residual: torch.Tensor | None,
forward_batch: ForwardBatch,
) -> tuple[torch.Tensor, torch.Tensor]:
if is_dp_attention_enabled():
@@ -669,11 +701,9 @@ class NemotronHAttentionDecoderLayer(NemotronHAttnLikeDecoderLayer):
)
return hidden_states, residual
if residual is None:
residual = hidden_states
hidden_states = self.norm(hidden_states)
else:
hidden_states, residual = self.norm(hidden_states, residual)
hidden_states, residual = input_norm_maybe_fuse_allreduce(
self.norm, hidden_states, residual
)
hidden_states = self.mixer.forward(
hidden_states=hidden_states, forward_batch=forward_batch
@@ -700,7 +730,7 @@ class NemotronHModel(nn.Module):
self,
*,
config: NemotronHConfig,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
):
super().__init__()
@@ -747,9 +777,9 @@ class NemotronHModel(nn.Module):
input_ids: torch.Tensor,
positions: torch.Tensor,
forward_batch: ForwardBatch,
pp_proxy_tensors: Optional[PPProxyTensors] = None,
inputs_embeds: Optional[torch.Tensor] = None,
) -> Union[torch.Tensor, PPProxyTensors]:
pp_proxy_tensors: PPProxyTensors | None = None,
inputs_embeds: torch.Tensor | None = None,
) -> torch.Tensor | PPProxyTensors:
if self.pp_group.is_first_rank:
if inputs_embeds is not None:
hidden_states = inputs_embeds
@@ -819,7 +849,7 @@ class NemotronHForCausalLM(nn.Module):
self,
*,
config: NemotronHConfig,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
):
super().__init__()
@@ -874,7 +904,7 @@ class NemotronHForCausalLM(nn.Module):
def _init_model(
self,
config: NemotronHConfig,
quant_config: Optional[QuantizationConfig] = None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
):
return NemotronHModel(
@@ -984,8 +1014,8 @@ class NemotronHForCausalLM(nn.Module):
input_ids: torch.Tensor,
positions: torch.Tensor,
forward_batch: ForwardBatch,
input_embeds: Optional[torch.Tensor] = None,
pp_proxy_tensors: Optional[PPProxyTensors] = None,
input_embeds: torch.Tensor | None = None,
pp_proxy_tensors: PPProxyTensors | None = None,
):
hidden_states = self.model.forward(
input_ids, positions, forward_batch, pp_proxy_tensors, input_embeds
@@ -147,6 +147,7 @@ class NemotronHMTPMoEDecoderLayer(NemotronHMoEDecoderLayer):
self.prev_layer_is_attn = layer_idx > 0 and is_attn_layer(
_pat[(layer_idx - 1) % len(_pat)]
)
self.layer_communicator.is_last_layer = True
if has_start_projections:
self.enorm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
+32 -1
View File
@@ -4,10 +4,12 @@ import torch
from torch import nn
from sglang.srt.configs.nemotron_h import ATTENTION, MAMBA
from sglang.srt.distributed import tensor_model_parallel_all_reduce
from sglang.srt.layers.communicator import (
LayerCommunicator,
LayerScatterModes,
ScatterMode,
apply_flashinfer_allreduce_fusion,
)
from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
@@ -57,7 +59,11 @@ def _build_layer_scatter_modes() -> LayerScatterModes:
def make_layer_communicator(
layer_norm: RMSNorm, *, for_attn: bool, allow_reduce_scatter: bool = False
layer_norm: RMSNorm,
*,
for_attn: bool,
allow_reduce_scatter: bool = False,
is_last_layer: bool = False,
) -> LayerCommunicator:
return LayerCommunicator(
layer_scatter_modes=_build_layer_scatter_modes(),
@@ -65,4 +71,29 @@ def make_layer_communicator(
post_attention_layernorm=nn.Identity() if for_attn else layer_norm,
force_layernorm_before_dp_gather=True,
allow_reduce_scatter=allow_reduce_scatter,
is_last_layer=is_last_layer,
)
def input_norm_maybe_fuse_allreduce(
norm: RMSNorm,
hidden_states: torch.Tensor,
residual: torch.Tensor | None,
) -> tuple[torch.Tensor, torch.Tensor]:
if residual is not None and getattr(
hidden_states, "_sglang_needs_allreduce_fusion", False
):
if apply_flashinfer_allreduce_fusion(hidden_states.shape[0]) and hasattr(
norm, "forward_with_allreduce_fusion"
):
return norm.forward_with_allreduce_fusion(
hidden_states, residual, use_attn_tp_group=False
)
hidden_states = tensor_model_parallel_all_reduce(hidden_states)
return norm(hidden_states, residual)
if residual is None:
residual = hidden_states
hidden_states = norm(hidden_states)
return hidden_states, residual
return norm(hidden_states, residual)
+2
View File
@@ -2780,6 +2780,8 @@ class ServerArgs:
"Qwen3_5MoeForConditionalGeneration",
"InternS2PreviewForConditionalGeneration",
"Qwen3_5ForConditionalGeneration",
"NemotronHForCausalLM",
"NemotronHPuzzleForCausalLM",
]
and is_sm100_supported()
and self.tp_size > 1