From 5d721884a324d7b7d46cadc71da5d8e150eb608b Mon Sep 17 00:00:00 2001 From: Cheng Wan Date: Sat, 26 Sep 2026 09:27:26 -0700 Subject: [PATCH] Nemotron-H: build each layer's boundaries from its stage and the previous one A Nemotron-H layer is a single stage: a Mamba / attention mixer, or an MLP / MoE FFN. Which stage completed a reduction was decided in the model from the layer types on both sides (input_is_reduced, feeds_mlp_layer, skip_reduce). An FFN whose input was already complete folded it into the residual in bf16 and handed prepare_mlp a zero update to reduce. The aux snapshot reduced a copy after a pattern check, and the MTP MoE layer was marked last after construction. - nemotron_h_utils declares each stage. A mixer computes on the attention's rows and owes its attention-TP sum to an FFN stage after it, always, or to a mixer after it when the fused kernel takes it. An FFN computes on the TP group's rows (a MoE dispatched by an a2a backend on this rank's own rows, as the scatter modes had it), and may leave its sum to a mixer after it. boundary_layout.stage_edges builds a layer's two boundaries from its own declaration and the previous stage's output. LayerCommunicator takes them as a LayerStage and runs them through make_boundary. - LayerCommunicator.mixer_exit, the mixer's counterpart of ffn_exit, decides once from what the stage's output declares whether the mixer skips its all-reduce. The FFN stages run through ffn_exit. - An FFN whose input is complete adds it and normalizes. At the start of the layer stack it reads its input with no residual (AddAndNorm.update_and_read_ffn_input with residual None). An FFN on each rank's own slice (a MoE dispatched by an a2a backend) first takes this rank's slice of the complete input and of the residual. - The layer loop's snapshot reads from the boundary's declaration whether the value there owes an attention-TP sum. - The MTP MoE layer takes its last-layer fact from the MTP pattern it is built with; the assignment after construction is removed. - The guard that no FFN behind ffn_exit reduces over attention TP also checks subclasses of ffn_exit users and the FFN a stage holds in `mixer`. Before, it did not see Nemotron-H. - token_axis_sizes becomes public for the model's declarations. On the transitions real checkpoints use (mixer -> FFN, FFN -> mixer, mixer -> mixer), each boundary runs the same reductions on the same tensors. Two differ: - An FFN -> mixer all-reduce that the fused kernel does not take may run in the next stage's input instead of in the FFN: the same all-reduce on the same group and tensor, as for the other models on ffn_exit. - At attention TP 1 the snapshot no longer all-reduces over the size-1 group. On patterns with an FFN after an FFN or at the start of the stack, the all-reduce of the zero update is gone. The add now happens inside the fused add + norm instead of in bf16 before it. --- python/sglang/srt/layers/boundary_layout.py | 44 ++++- python/sglang/srt/layers/communicator.py | 159 ++++++++++++++-- python/sglang/srt/models/nemotron_h.py | 120 ++++-------- python/sglang/srt/models/nemotron_h_mtp.py | 1 - python/sglang/srt/models/nemotron_h_utils.py | 110 +++++++++-- .../unit/layers/test_communicator_ffn_exit.py | 38 +++- .../layers/test_declared_decoder_boundary.py | 26 +++ .../models/test_nemotron_h_aux_capture.py | 23 ++- .../models/test_nemotron_h_mtp_reduction.py | 4 +- .../unit/models/test_nemotron_h_stages.py | 171 ++++++++++++++++++ 10 files changed, 560 insertions(+), 136 deletions(-) create mode 100644 test/registered/unit/models/test_nemotron_h_stages.py diff --git a/python/sglang/srt/layers/boundary_layout.py b/python/sglang/srt/layers/boundary_layout.py index 8b30f64fd105..ff7df0f1fad4 100644 --- a/python/sglang/srt/layers/boundary_layout.py +++ b/python/sglang/srt/layers/boundary_layout.py @@ -14,7 +14,7 @@ """Token layouts of the tensors handed across layer communication boundaries.""" from enum import Enum, auto -from typing import FrozenSet, Mapping, Optional +from typing import FrozenSet, Mapping, Optional, Tuple import msgspec @@ -108,6 +108,14 @@ class DecoderLayerSides(msgspec.Struct, frozen=True): residual_joins_attention_sum: bool = False +class StageDecl(msgspec.Struct, frozen=True): + """A computing stage's two sides: the rows its input must be on, and what + its output is.""" + + input: StageInput + output: StageOutput + + class EdgeDecl(msgspec.Struct, frozen=True): """One boundary between two stages, as the layer that runs one side of it sees it: what arrives from the producer, what the consumer needs, and the @@ -170,6 +178,40 @@ def decoder_layer_edges(sides: DecoderLayerSides) -> DecoderLayerEdges: ) +def stage_edges( + *, previous: Optional[StageOutput], stage: StageDecl, rows: Layout +) -> Tuple[EdgeDecl, EdgeDecl]: + """The two boundaries of a layer that is one stage of a sequence of stages: + into it from the previous stage's output (None at the start of the layer + stack) and out of it onto ``rows``, the rows every layer hands on and the + residual is on between stages. The previous output arrives on those rows: a + stage whose output is elsewhere (an FFN on the TP group) moves it back there + first, and what it may leave of its sum comes with the value; otherwise the + value is complete. The residual follows the input onto a finer slice and + stays where it is when the input is gathered.""" + arrived = ( + StageOutput( + rows, + group=previous.group, + always_leaves=previous.always_leaves, + leaves_for_next_layer=previous.leaves_for_next_layer, + ) + if previous is not None + and (previous.always_leaves or previous.leaves_for_next_layer) + else StageOutput(rows) + ) + during = stage.input.layout if rows.sharded <= stage.input.layout.sharded else rows + return ( + EdgeDecl(produced=arrived, need=stage.input, residual=rows, residual_to=during), + EdgeDecl( + produced=stage.output, + need=StageInput(rows), + residual=during, + residual_to=rows, + ), + ) + + def sequence_parallel_layer_sides( *, axis_sizes: Mapping[TokenAxis, int] ) -> DecoderLayerSides: diff --git a/python/sglang/srt/layers/communicator.py b/python/sglang/srt/layers/communicator.py index 2412046c16b5..eca837e61c78 100644 --- a/python/sglang/srt/layers/communicator.py +++ b/python/sglang/srt/layers/communicator.py @@ -725,6 +725,11 @@ def update_and_read_attention_input( ) def update_and_read_ffn_input(self, hidden_states, residual, norm): + if residual is None: + # The layer stack starts at this FFN: its input is the residual. + if hidden_states.shape[0] == 0: + return hidden_states, hidden_states + return norm(hidden_states), hidden_states if hidden_states.shape[0] == 0: return hidden_states, residual return norm(hidden_states, residual) @@ -820,6 +825,9 @@ def __init__( allow_deferred_ffn_reduction: bool = True, # How the layer writes its residual and reads its stages' inputs. residual_ops: ResidualOps = ADD_AND_NORM, + # A layer that is one stage of a sequence of stages, instead of an + # attention followed by an FFN. + stage: Optional["LayerStage"] = None, ): self.layer_scatter_modes = layer_scatter_modes self.input_layernorm = input_layernorm @@ -839,6 +847,18 @@ def __init__( ) # The fused kernels every batch's attention input tries first. self._attn_input_fusions = self._select_attn_input_fusions() + self._speculative_algo = SpeculativeAlgorithm.from_string( + get_spec().speculative_algorithm + ) + # LoRA kernels need the per-layer token layout only under DP attention. + self._publish_lora_layout = get_parallel().enable_dp_attention and bool( + get_lora().enable_lora + ) + if stage is not None: + self._init_stage(stage) + return + # Its two boundaries, for a layer that is one stage. + self.stage_edges = None # The steps the layer's ordinary batches run. sides = self._declared_sides() self._declared = sides @@ -870,20 +890,13 @@ def __init__( if sides is not None and self._input_can_be_scattered() else None ) - self._speculative_algo = SpeculativeAlgorithm.from_string( - get_spec().speculative_algorithm - ) - # LoRA kernels need the per-layer token layout only under DP attention. - self._publish_lora_layout = get_parallel().enable_dp_attention and bool( - get_lora().enable_lora - ) # Under LayerNorm SP, the steps the layer runs while the region is # active; None without SP. The two are exclusive: SP runs a model without # q_lora, which input-scattered attention needs. self._sp_steps = ( _select_boundary_steps( - sequence_parallel_layer_sides(axis_sizes=_token_axis_sizes()), + sequence_parallel_layer_sides(axis_sizes=token_axis_sizes()), residual_ops=residual_ops, attention_fusions=self._attn_input_fusions, enters_stack=self.layer_scatter_modes.is_first_layer, @@ -892,12 +905,51 @@ def __init__( else None ) + def _init_stage(self, stage: "LayerStage") -> None: + """A layer that is one stage: every batch runs the two boundaries its + declarations give.""" + if get_parallel().attn_cp_size > 1 or layernorm_sp.layernorm_sp_enabled(): + raise NotImplementedError( + "a layer that is one stage with attention CP or LayerNorm SP" + ) + self._declared = None + self._cp_steps = self._input_scattered_steps = self._sp_steps = None + self.stage_edges = stage.edges + into_edge, out_edge = stage.edges + reads_ffn = stage.reads is InputRead.FFN + into = make_boundary( + into_edge, + reads=stage.reads, + fusions=( + self._select_mlp_input_fusions() + if reads_ffn + else self._attn_input_fusions + ), + force_layernorm_before_gather=self.force_layernorm_before_dp_gather, + residual_ops=self._residual_ops, + enters_stack=stage.enters_stack, + ) + out = make_boundary(out_edge, reads=None, residual_ops=self._residual_ops) + self._steps = BoundarySteps( + attention_prepare=_another_stage if reads_ffn else into.prepare, + attention_input=_another_stage if reads_ffn else into.input_move, + ffn_input=into.prepare if reads_ffn else _another_stage, + ffn_input_rows=into.input_rows, + ffn_output=out_edge.produced, + ffn_output_move=out.output_move, + ffn_output_move_completes_sum=out.output_move_completes_sum, + ffn_sum_is_movable=out_edge.produced.group is not None, + fused=into.fused, + ) + @property def input_rows(self) -> Layout: """The rows the layer's input and residual arrive on in a batch that runs its ordinary steps.""" if self._declared is not None: return self._declared.input_rows + if self.stage_edges is not None: + return self.stage_edges[0].residual # Such a batch holds every token on each CP rank. return scatter_mode_layouts( attn_dp_size=self._context.attn_dp_size, @@ -974,7 +1026,7 @@ def gathers_for_tbo(sparse: bool, next_sparse: Optional[bool]) -> bool: # Otherwise under CP the FFN completes its own sum. may_leave = may_leave_to_reduce_scatter = parallel.attn_cp_size == 1 return decoder_layer_sides( - axis_sizes=_token_axis_sizes(cp_active=cp_active), + axis_sizes=token_axis_sizes(cp_active=cp_active), ffn_on_local_rows=on_local_rows(modes.is_layer_sparse), previous_on_local_rows=( not modes.is_first_layer @@ -1054,14 +1106,14 @@ def _steps_for_input_scattered(self, sides: DecoderLayerSides) -> "BoundarySteps write-back is not a plain add stays on each rank's slice.""" if self._residual_ops.adds_plainly: scattered = input_scattered_layer_sides( - axis_sizes=_token_axis_sizes(), + axis_sizes=token_axis_sizes(), ffn_group=sides.ffn_output.group, hands_on_partial=self.allow_reduce_scatter and not self.is_last_layer, ) handoff = _hand_qkv_hook_its_input else: scattered = scattered_residual_layer_sides( - axis_sizes=_token_axis_sizes(), + axis_sizes=token_axis_sizes(), ffn_group=sides.ffn_output.group, is_first_layer=self.layer_scatter_modes.is_first_layer, is_last_layer=self.is_last_layer, @@ -1564,6 +1616,12 @@ def _complete_ffn_output_now( residual = None return hidden_states, residual + def mixer_exit(self, forward_batch: ForwardBatch) -> "MixerExit": + """Decide once whether this stage's mixer (an attention-like stage) + skips its output all-reduce. Use the result as a context manager around + the mixer, then call ``finish``.""" + return MixerExit(self, forward_batch) + def ffn_exit(self, forward_batch: ForwardBatch) -> "FfnExit": """Decide once how this layer's FFN output reduction completes. Use the result as a context manager around the FFN call, then call ``finish``.""" @@ -1722,6 +1780,42 @@ def _leave_to_next_layer( return wrap(hidden_states), residual +class MixerExit: + """The scope that publishes a mixer's decision while it runs: inside the + ``with`` block ``fuse_mlp_allreduce`` on ``get_forward()`` tells its + row-parallel output projection to skip the all-reduce. It skips when the + stage's output always leaves its sum (to an FFN stage, which completes it in + its input), and when it may leave it and the fused kernel takes it into the + next input norm.""" + + __slots__ = ("skips_reduction", "_hands_on", "_scope") + + def __init__(self, communicator: LayerCommunicator, forward_batch: ForwardBatch): + produced = communicator._batch_steps(forward_batch).ffn_output + self._hands_on = ( + produced.leaves_for_next_layer + and communicator.should_fuse_mlp_allreduce_with_next_layer(forward_batch) + ) + self.skips_reduction = produced.always_leaves or self._hands_on + self._scope = get_forward().scoped(fuse_mlp_allreduce=self.skips_reduction) + + def __enter__(self) -> "MixerExit": + self._scope.__enter__() + return self + + def __exit__(self, *exc_info): + return self._scope.__exit__(*exc_info) + + def finish( + self, hidden_states: torch.Tensor + ) -> Union[torch.Tensor, UnreducedOutput]: + """The mixer's output: its partial sum to the FFN stage after it, as a + value that owes the sum to a mixer after it, or complete.""" + if self._hands_on: + return UnreducedOutput(hidden_states, group=get_parallel().tp_group) + return hidden_states + + class FfnExit: """The scope that publishes an FfnCompletion while the FFN runs: inside the ``with`` block it is ``fuse_mlp_allreduce`` / ``mlp_reduce_scatter`` / @@ -2135,7 +2229,7 @@ def _sum_group(group: SumGroup) -> GroupCoordinator: return post_experts_reduction_group() -def _token_axis_sizes(*, cp_active: bool = False) -> Dict[TokenAxis, int]: +def token_axis_sizes(*, cp_active: bool = False) -> Dict[TokenAxis, int]: """The token axes' sizes for a batch: attention CP shards tokens only on a CP extend (``cp_active``); otherwise every CP rank holds them all.""" parallel = get_parallel() @@ -2151,7 +2245,7 @@ def tbo_split_moves(layer_input_rows: Layout) -> Tuple[Callable, Callable]: rows: from the rows the first overlapped layer takes to the attention's, and back again for each half.""" attention = Layout.sharded_over( - TokenAxis.ATTN_DP, TokenAxis.ATTN_CP, axis_sizes=_token_axis_sizes() + TokenAxis.ATTN_DP, TokenAxis.ATTN_CP, axis_sizes=token_axis_sizes() ) pair = CommunicateSummableTensorPairFn if layer_input_rows == attention: @@ -2353,6 +2447,11 @@ def _attention_input_step( ) +def _another_stage(*args, **kwargs): + """The steps of a stage a single-stage layer does not have.""" + raise RuntimeError("this layer is one stage and does not have the other") + + class InputRead(Enum): """How a boundary's consumer reads its input from the residual: with the attention input norm (prepare_attn) or with the FFN input norm and its @@ -2362,6 +2461,18 @@ class InputRead(Enum): FFN = auto() +class LayerStage(msgspec.Struct, frozen=True): + """A layer that is one stage of a sequence of stages, each an + attention-like mixer or an FFN: how it reads its input, its two boundaries + (into it, and out of it onto the rows every layer hands on, as + ``boundary_layout.stage_edges`` gives them), and whether the layer stack + starts at it.""" + + reads: InputRead + edges: Tuple[EdgeDecl, EdgeDecl] + enters_stack: bool = False + + class Boundary(msgspec.Struct, frozen=True): """The steps one layer runs at one boundary, chosen from both sides' declarations. A layer runs the consumer's half of a boundary into one of @@ -2595,18 +2706,18 @@ def _select_ffn_input( sliced = need.layout.sharded - produced.layout.sharded if sliced: # Each attention-TP rank takes its own slice: the reduce-scatter - # completes the attention-TP sum and slices in one collective. + # completes the attention-TP sum and slices in one collective; a + # complete value is only sliced. if ( sliced != {TokenAxis.ATTN_TP_SCATTER} or gathered - or not owes_attention_tp or residual_to != need.layout or residual not in (produced.layout, need.layout) ): raise NotImplementedError(f"{produced=} {residual=} {need=}") return ( partial( - _mlp_input_scatter, + _mlp_input_scatter if owes_attention_tp else _mlp_input_slice, scatters_residual=residual != residual_to, residual_ops=residual_ops, ), @@ -2878,6 +2989,22 @@ def _mlp_input_scatter( return residual_ops.update_and_read_ffn_input(hidden_states, residual, layernorm) +def _mlp_input_slice( + hidden_states: torch.Tensor, + residual: Optional[torch.Tensor], + forward_batch: ForwardBatch, + layernorm: torch.nn.Module, + context: CommunicateContext, + *, + scatters_residual: bool, + residual_ops: ResidualOps = ADD_AND_NORM, +): + hidden_states = _redistribute_to_attn_tp_shards(hidden_states, context).clone() + if scatters_residual and residual is not None: + residual = residual_ops.residual_to_attn_tp_shard(residual, context) + return residual_ops.update_and_read_ffn_input(hidden_states, residual, layernorm) + + def _mlp_input_on_residual_shard( hidden_states: torch.Tensor, residual: torch.Tensor, diff --git a/python/sglang/srt/models/nemotron_h.py b/python/sglang/srt/models/nemotron_h.py index 983f393a4c75..ae6b2ae6abd9 100644 --- a/python/sglang/srt/models/nemotron_h.py +++ b/python/sglang/srt/models/nemotron_h.py @@ -78,11 +78,7 @@ replace_prefix, replace_substrings, ) -from sglang.srt.models.nemotron_h_utils import ( - feeds_mlp_layer, - is_attn_layer, - make_layer_communicator, -) +from sglang.srt.models.nemotron_h_utils import make_layer_communicator from sglang.srt.models.utils import WeightsMapper from sglang.srt.runtime_context import get_exec, get_forward, get_parallel from sglang.srt.utils import ( @@ -376,66 +372,28 @@ def forward( class NemotronHMLPLikeDecoderLayer(nn.Module): - """MLP half of a decoder layer. A preceding Mamba / attention layer hands over - its unreduced output, as between the two halves of a standard decoder layer.""" + """An FFN (MLP / MoE) stage. Its input is the previous stage's output: a + mixer's attention partial sum, or a value that is complete or carries the + sum an FFN before it left.""" - def _init_layer_communicator( - self, config: NemotronHConfig, layer_idx: int, *, is_sparse: bool - ) -> None: - pattern = config.hybrid_override_pattern - self.input_is_reduced = layer_idx == 0 or not is_attn_layer( - pattern[layer_idx - 1] - ) - self.feeds_mlp_layer = feeds_mlp_layer(pattern, layer_idx) + def _init_layer_communicator(self, config: NemotronHConfig, layer_idx: int): self.layer_communicator = make_layer_communicator( - self.norm, - for_attn=False, - allow_reduce_scatter=True, - is_sparse=is_sparse, - is_last_layer=layer_idx == len(pattern) - 1, + self.norm, pattern=config.hybrid_override_pattern, layer_idx=layer_idx ) def forward( self, *, - hidden_states: torch.Tensor, + hidden_states: torch.Tensor | UnreducedOutput, residual: torch.Tensor | None, forward_batch: ForwardBatch, - ) -> tuple[torch.Tensor, torch.Tensor]: - if self.input_is_reduced: - # Fold the reduced input into the residual, leaving prepare_mlp a - # zero update to reduce. - residual = hidden_states if residual is None else hidden_states + residual - hidden_states = torch.zeros_like(hidden_states) + ) -> tuple[torch.Tensor | UnreducedOutput, torch.Tensor]: hidden_states, residual = self.layer_communicator.prepare_mlp( hidden_states, residual, forward_batch ) - mlp_reduce_scatter = self.layer_communicator.should_use_reduce_scatter( - forward_batch - ) - # prepare_mlp ignores a deferred reduction, so only defer into a - # Mamba / attention layer. - fuse_mlp_allreduce = ( - not self.feeds_mlp_layer - and self.layer_communicator.should_fuse_mlp_allreduce_with_next_layer( - forward_batch - ) - ) - with get_forward().scoped( - fuse_mlp_allreduce=fuse_mlp_allreduce, - mlp_reduce_scatter=mlp_reduce_scatter, - ): + with self.layer_communicator.ffn_exit(forward_batch) as ffn_exit: hidden_states = self.mixer.forward(hidden_states) - if fuse_mlp_allreduce: - hidden_states = UnreducedOutput( - hidden_states, - group=self.layer_communicator.ffn_reduction_group(forward_batch), - ) - else: - hidden_states, residual = self.layer_communicator.postprocess_layer( - hidden_states, residual, forward_batch - ) - return hidden_states, residual + return ffn_exit.finish(hidden_states, residual) class NemotronHMLPDecoderLayer(NemotronHMLPLikeDecoderLayer): @@ -469,7 +427,7 @@ def __init__( ) self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon) - self._init_layer_communicator(config, layer_idx, is_sparse=False) + self._init_layer_communicator(config, layer_idx) class NemotronHMoEDecoderLayer(NemotronHMLPLikeDecoderLayer): @@ -492,52 +450,37 @@ def __init__( ) self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon) - self._init_layer_communicator(config, layer_idx, is_sparse=True) + self._init_layer_communicator(config, layer_idx) class NemotronHAttnLikeDecoderLayer(nn.Module): - """Attention half of a decoder layer. Before an MLP layer the mixer leaves its - output unreduced for prepare_mlp; otherwise the mixer reduces it itself.""" + """A mixer (Mamba / attention) stage. Before an FFN stage it leaves its + output all-reduce to that stage's input; otherwise it reduces the output + itself unless the fused kernel takes the sum into the next input norm.""" def _init_layer_communicator(self, config: NemotronHConfig, layer_idx: int): - self.feeds_mlp_layer = feeds_mlp_layer( - config.hybrid_override_pattern, layer_idx - ) self.layer_communicator = make_layer_communicator( - self.norm, - for_attn=True, - is_last_layer=layer_idx == len(config.hybrid_override_pattern) - 1, + self.norm, pattern=config.hybrid_override_pattern, layer_idx=layer_idx ) def forward( self, *, - hidden_states: torch.Tensor, + hidden_states: torch.Tensor | UnreducedOutput, residual: torch.Tensor | None, forward_batch: ForwardBatch, - ) -> tuple[torch.Tensor, torch.Tensor]: + ) -> tuple[torch.Tensor | UnreducedOutput, torch.Tensor]: hidden_states, residual = self.layer_communicator.prepare_attn( hidden_states, residual, forward_batch ) if forward_batch.forward_mode.is_idle(): return hidden_states, residual - fuse_mlp_allreduce = ( - not self.feeds_mlp_layer - and self.layer_communicator.should_fuse_mlp_allreduce_with_next_layer( - forward_batch - ) - ) - skip_reduce = self.feeds_mlp_layer or fuse_mlp_allreduce - with get_forward().scoped(fuse_mlp_allreduce=skip_reduce): + with self.layer_communicator.mixer_exit(forward_batch) as mixer_exit: hidden_states = self._forward_mixer( - hidden_states, forward_batch, skip_reduce - ) - if fuse_mlp_allreduce: - hidden_states = UnreducedOutput( - hidden_states, group=get_parallel().tp_group + hidden_states, forward_batch, mixer_exit.skips_reduction ) - return hidden_states, residual + return mixer_exit.finish(hidden_states), residual class NemotronHMambaDecoderLayer(NemotronHAttnLikeDecoderLayer): @@ -779,18 +722,23 @@ def get_layer(idx: int, prefix: str): self.layers_to_capture: set[int] = set() def _capture_hidden_states(self, hidden_states, residual, boundary_idx): - pattern = self.config.hybrid_override_pattern - if ( - residual is not None - and boundary_idx > 0 - and is_attn_layer(pattern[boundary_idx - 1]) - and feeds_mlp_layer(pattern, boundary_idx - 1) - ): - # Reduce a copy so the MLP layer still receives a TP partial. + if residual is not None and self._owes_attention_sum_at(boundary_idx): + # Reduce a copy so the FFN stage still receives its partial sum. hidden_states = attn_tp_all_reduce(hidden_states.clone()) # Later norms update the input residual in place. return hidden_states.clone() if residual is None else hidden_states + residual + def _owes_attention_sum_at(self, boundary_idx: int) -> bool: + """Whether the value at this boundary is a mixer's partial sum that the + FFN stage after it completes: what the layer after the boundary + declares it takes, or at this rank's end what the last layer declares + it hands on.""" + if boundary_idx < self.end_layer: + edge = self.layers[boundary_idx].layer_communicator.stage_edges[0] + else: + edge = self.layers[boundary_idx - 1].layer_communicator.stage_edges[1] + return edge.produced.always_leaves + def forward( self, input_ids: torch.Tensor, diff --git a/python/sglang/srt/models/nemotron_h_mtp.py b/python/sglang/srt/models/nemotron_h_mtp.py index 6522fc06c6cb..09bc6e7809bf 100644 --- a/python/sglang/srt/models/nemotron_h_mtp.py +++ b/python/sglang/srt/models/nemotron_h_mtp.py @@ -141,7 +141,6 @@ def __init__( ) self.has_start_projections = has_start_projections self.has_end_norm = has_end_norm - self.layer_communicator.is_last_layer = True if has_start_projections: self.enorm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon) diff --git a/python/sglang/srt/models/nemotron_h_utils.py b/python/sglang/srt/models/nemotron_h_utils.py index 6c89fe22ec4b..77f142ce2f13 100644 --- a/python/sglang/srt/models/nemotron_h_utils.py +++ b/python/sglang/srt/models/nemotron_h_utils.py @@ -1,12 +1,24 @@ """Layer-communication helpers for the Nemotron-H model.""" -from torch import nn +from typing import Optional -from sglang.srt.configs.nemotron_h import ATTENTION, MAMBA +from sglang.srt.configs.nemotron_h import ATTENTION, MAMBA, MOE +from sglang.srt.layers.boundary_layout import ( + Layout, + StageDecl, + StageInput, + StageOutput, + SumGroup, + TokenAxis, + stage_edges, +) from sglang.srt.layers.communicator import ( + InputRead, LayerCommunicator, LayerScatterModes, + LayerStage, ScatterMode, + token_axis_sizes, ) from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.moe.utils import get_moe_a2a_backend @@ -18,9 +30,75 @@ def is_attn_layer(layer_type: str) -> bool: return layer_type in ATTN_LAYERS -def feeds_mlp_layer(pattern: str, layer_idx: int) -> bool: - next_idx = layer_idx + 1 - return next_idx < len(pattern) and not is_attn_layer(pattern[next_idx]) +def _stage_kind(pattern: str, layer_idx: int) -> Optional[InputRead]: + """How the stage at ``layer_idx`` reads its input: a Mamba or attention + mixer like an attention, an MLP or MoE like an FFN; None past either end.""" + if not 0 <= layer_idx < len(pattern): + return None + return InputRead.ATTENTION if is_attn_layer(pattern[layer_idx]) else InputRead.FFN + + +def _stage_decl(pattern: str, layer_idx: int) -> StageDecl: + """What the stage at ``layer_idx`` declares. A mixer computes on the + attention's rows and leaves its attention-TP sum to an FFN stage after it, + or, when a fused kernel takes it, to a mixer after it. An FFN computes on + the TP group's rows, a MoE dispatched by an a2a backend on this rank's own, + and may leave its sum to a mixer after it.""" + axis_sizes = token_axis_sizes() + attention = Layout.sharded_over( + TokenAxis.ATTN_DP, TokenAxis.ATTN_CP, axis_sizes=axis_sizes + ) + following = _stage_kind(pattern, layer_idx + 1) + if is_attn_layer(pattern[layer_idx]): + owes = axis_sizes[TokenAxis.ATTN_TP_SCATTER] > 1 + return StageDecl( + StageInput(attention), + StageOutput( + attention, + group=SumGroup.ATTN_TP if owes else None, + always_leaves=owes and following is InputRead.FFN, + leaves_for_next_layer=owes and following is InputRead.ATTENTION, + ), + ) + sparse = pattern[layer_idx] == MOE + if sparse and not get_moe_a2a_backend().is_none(): + local = Layout.sharded_over( + TokenAxis.ATTN_DP, + TokenAxis.ATTN_CP, + TokenAxis.ATTN_TP_SCATTER, + axis_sizes=axis_sizes, + ) + return StageDecl(StageInput(local), StageOutput(local)) + full = Layout.sharded_over(axis_sizes=axis_sizes) + return StageDecl( + StageInput(full), + StageOutput( + full, + group=SumGroup.MOE_OUTPUT if sparse else SumGroup.TP, + leaves_for_next_layer=following is InputRead.ATTENTION, + leaves_for_reduce_scatter=True, + leaves_for_reduce_scatterv=True, + ), + ) + + +def layer_stage(pattern: str, layer_idx: int) -> LayerStage: + """The layer's stage and its two boundaries, from its own declaration and + the previous stage's.""" + rows = Layout.sharded_over( + TokenAxis.ATTN_DP, TokenAxis.ATTN_CP, axis_sizes=token_axis_sizes() + ) + return LayerStage( + reads=_stage_kind(pattern, layer_idx), + edges=stage_edges( + previous=( + _stage_decl(pattern, layer_idx - 1).output if layer_idx > 0 else None + ), + stage=_stage_decl(pattern, layer_idx), + rows=rows, + ), + enters_stack=layer_idx == 0, + ) def _build_layer_scatter_modes( @@ -43,19 +121,21 @@ def _build_layer_scatter_modes( def make_layer_communicator( - layer_norm: RMSNorm, - *, - for_attn: bool, - allow_reduce_scatter: bool = False, - is_sparse: bool = False, - is_last_layer: bool = False, + layer_norm: RMSNorm, *, pattern: str, layer_idx: int ) -> LayerCommunicator: + """The communicator of a layer that is one stage: only its own norm, and + boundaries built from the stages next to it in the pattern.""" + stage = layer_stage(pattern, layer_idx) + for_attn = stage.reads is InputRead.ATTENTION return LayerCommunicator( - layer_scatter_modes=_build_layer_scatter_modes(is_sparse, is_last_layer), - input_layernorm=layer_norm if for_attn else nn.Identity(), - post_attention_layernorm=nn.Identity() if for_attn else layer_norm, + layer_scatter_modes=_build_layer_scatter_modes( + pattern[layer_idx] == MOE, is_last_layer=layer_idx == len(pattern) - 1 + ), + input_layernorm=layer_norm if for_attn else None, + post_attention_layernorm=None if for_attn else layer_norm, # With attention TP > 1, the default gather adds the residual to one # rank's partial in bf16 before the cross-rank sum. force_layernorm_before_dp_gather=True, - allow_reduce_scatter=allow_reduce_scatter, + allow_reduce_scatter=not for_attn, + stage=stage, ) diff --git a/test/registered/unit/layers/test_communicator_ffn_exit.py b/test/registered/unit/layers/test_communicator_ffn_exit.py index 71b324df3146..71b14b11deda 100644 --- a/test/registered/unit/layers/test_communicator_ffn_exit.py +++ b/test/registered/unit/layers/test_communicator_ffn_exit.py @@ -426,8 +426,9 @@ def test_the_all_reduce_runs_before_the_scatter(self): SRT_DIR = Path(sglang.__file__).resolve().parent / "srt" -# Attributes a decoder layer holds its FFN, or parts of it, in. -FFN_ATTRS = {"mlp", "moe", "shared_expert", "shared_experts", "share_expert"} +# Attributes a decoder layer holds its FFN, or parts of it, in. A layer that is +# one FFN stage holds it in ``mixer``. +FFN_ATTRS = {"mlp", "moe", "shared_expert", "shared_experts", "share_expert", "mixer"} def can_be_true(value, params): @@ -498,6 +499,27 @@ def attention_tp_reductions(cls): return found +def ffn_exit_classes(tree, source): + """Classes in a module that call ffn_exit, and the classes there that inherit + from them (whose constructors build the FFN the inherited forward runs).""" + classes = [node for node in tree.body if isinstance(node, ast.ClassDef)] + users = { + cls.name + for cls in classes + if ".ffn_exit(" in ast.get_source_segment(source, cls) + } + grew = True + while grew: + grew = False + for cls in classes: + if cls.name not in users and any( + getattr(base, "id", None) in users for base in cls.bases + ): + users.add(cls.name) + grew = True + return [cls for cls in classes if cls.name in users] + + class TestFfnExitUsersOweATpSum(CustomTestCase): """Under attention DP the next layer's input completes a deferred FFN sum with an all-reduce over TP. An FFN whose linear layers reduce over the @@ -514,11 +536,7 @@ def test_no_ffn_behind_ffn_exit_reduces_over_attention_tp(self): or ".ffn_exit(" not in source ): continue - for cls in index.tree(module).body: - if not isinstance(cls, ast.ClassDef): - continue - if ".ffn_exit(" not in ast.get_source_segment(source, cls): - continue + for cls in ffn_exit_classes(index.tree(module), source): layers.append(f"{module}.{cls.name}") todo = [] for attr, name, call in submodule_constructions(cls): @@ -546,6 +564,12 @@ def test_no_ffn_behind_ffn_exit_reduces_over_attention_tp(self): ] todo += [(where, n) for _, n, _ in submodule_constructions(ffn)] self.assertTrue(layers and ffn_classes, "no FFN behind ffn_exit found") + # A layer that is one FFN stage inherits its forward and holds the FFN + # in ``mixer``. + nemotron = "sglang.srt.models.nemotron_h" + self.assertLessEqual( + {(nemotron, "NemotronHMLP"), (nemotron, "NemotronHMoE")}, ffn_classes + ) self.assertEqual(sorted(set(offenders)), []) diff --git a/test/registered/unit/layers/test_declared_decoder_boundary.py b/test/registered/unit/layers/test_declared_decoder_boundary.py index a5afabc566f3..a059cd6528ff 100644 --- a/test/registered/unit/layers/test_declared_decoder_boundary.py +++ b/test/registered/unit/layers/test_declared_decoder_boundary.py @@ -1363,6 +1363,32 @@ def test_the_order_of_a_residual_on_the_slice_is_declared(self): ) self.assertIs(getattr(steps, "func", steps), step) + def test_a_complete_value_is_sliced_onto_each_rank(self): + # An FFN on each rank's slice after an FFN, or at the start of the + # layer stack, takes a complete input. + sizes = self.SIZES + attention = comm.Layout.sharded_over(axis_sizes=sizes) + local = comm.Layout(frozenset({TokenAxis.ATTN_TP_SCATTER})) + step, _ = comm._select_ffn_input( + comm.StageOutput(attention), + residual=attention, + residual_to=local, + need=comm.StageInput(local), + force_layernorm_before_gather=False, + fusions=(), + ) + self.assertIs(step.func, comm._mlp_input_slice) + hidden = torch.arange(4.0)[:, None].expand(4, HIDDEN).clone() + residual = torch.ones(4, HIDDEN) + context = SimpleNamespace(attn_tp_size=2, attn_tp_rank=1) + out, out_residual = step(hidden, residual, None, Norm(), context) + torch.testing.assert_close(out_residual, hidden[2:] + 1) + torch.testing.assert_close(out, 2 * (hidden[2:] + 1)) + # At the start of the layer stack the input is the residual. + out, out_residual = step(hidden, None, None, Norm(), context) + torch.testing.assert_close(out_residual, hidden[2:]) + torch.testing.assert_close(out, 2 * hidden[2:]) + def test_which_layers_can_scatter_their_input(self): configured = dict(enable_attn_tp_input_scattered=True) for name, parallel, a2a, expected in ( diff --git a/test/registered/unit/models/test_nemotron_h_aux_capture.py b/test/registered/unit/models/test_nemotron_h_aux_capture.py index 9d0058f9ede7..7f8550fa38ef 100644 --- a/test/registered/unit/models/test_nemotron_h_aux_capture.py +++ b/test/registered/unit/models/test_nemotron_h_aux_capture.py @@ -56,12 +56,8 @@ def _build(pattern, tp, capture): layer = cls.__new__(cls) nn.Module.__init__(layer) layer.norm = _Norm() - if kind in "M*": - layer._init_layer_communicator(config, i) - layer.mixer = _Mixer(0.5, tp) - else: - layer._init_layer_communicator(config, i, is_sparse=False) - layer.mixer = _Mixer(0.25) + layer._init_layer_communicator(config, i) + layer.mixer = _Mixer(0.5 if kind in "M*" else 0.25, tp) if kind == "M": layer._forward_mamba = lambda h, batch, mixer=layer.mixer: mixer(h) layers.append(layer) @@ -73,7 +69,18 @@ class TestNemotronAuxCapture(CustomTestCase): def test_capture_reduces_only_its_snapshot(self): """Each auxiliary snapshot equals the full hidden state at its boundary, also under DP attention and after later norms update the residual in place.""" - for pattern in ("*-", "M-", "**-", "*", "*--", "-*"): + for pattern in ( + "*-", + "M-", + "**-", + "*", + "*--", + "-*", + "MM", + "-M", + "M-M*E", + "E*E-", + ): for dp_enabled, tp in ((True, 2), (True, 1), (False, 2)): with self.subTest(pattern=pattern, dp_enabled=dp_enabled, tp=tp): self._check(pattern, dp_enabled, tp) @@ -96,6 +103,8 @@ def reduce_in_place(x): get_flags().dp.override(enabled=dp_enabled), get_parallel().override( attn_tp_group=group, + tp_group=group, + moe_tp_group=group, launch_world_rank=0, tp_rank=0, tp_size=tp, diff --git a/test/registered/unit/models/test_nemotron_h_mtp_reduction.py b/test/registered/unit/models/test_nemotron_h_mtp_reduction.py index f8cef452a627..f6d6395d80b0 100644 --- a/test/registered/unit/models/test_nemotron_h_mtp_reduction.py +++ b/test/registered/unit/models/test_nemotron_h_mtp_reduction.py @@ -70,9 +70,7 @@ def test_attention_partial_is_reduced_once(self): layer.mixer = nn.Identity() layer.norm = _Norm() layer._init_layer_communicator( - SimpleNamespace(hybrid_override_pattern="*E"), - 1, - is_sparse=False, + SimpleNamespace(hybrid_override_pattern="*E"), 1 ) partial = torch.tensor([[1.0, 2.0], [3.0, 4.0]]) residual = torch.tensor([[7.0, 3.0], [5.0, 9.0]]) diff --git a/test/registered/unit/models/test_nemotron_h_stages.py b/test/registered/unit/models/test_nemotron_h_stages.py new file mode 100644 index 000000000000..2b8096efb738 --- /dev/null +++ b/test/registered/unit/models/test_nemotron_h_stages.py @@ -0,0 +1,171 @@ +import itertools +import unittest +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import torch + +from sglang.srt.layers.boundary_layout import ( + Layout, + StageDecl, + StageInput, + StageOutput, + SumGroup, + TokenAxis, + stage_edges, +) +from sglang.srt.layers.communicator import MixerExit, UnreducedOutput +from sglang.srt.layers.moe.utils import should_skip_mlp_all_reduce +from sglang.srt.models import nemotron_h_utils as utils +from sglang.srt.runtime_context import get_parallel +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=3, suite="base-a-test-cpu") + + +def sizes(*, dp=1, tp=1): + return { + TokenAxis.ATTN_DP: dp, + TokenAxis.ATTN_CP: 1, + TokenAxis.ATTN_TP_SCATTER: tp, + } + + +def stages(pattern, *, dp=1, tp=1, a2a=False): + backend = SimpleNamespace(is_none=lambda: not a2a) + with ( + patch.object(utils, "token_axis_sizes", return_value=sizes(dp=dp, tp=tp)), + patch.object(utils, "get_moe_a2a_backend", return_value=backend), + ): + return [utils.layer_stage(pattern, i) for i in range(len(pattern))] + + +class TestStageEdges(CustomTestCase): + """Each layer's two boundaries come from its own declaration and the + previous stage's; a boundary two layers share is seen the same way by both.""" + + PATTERNS = ("M-M*E", "MEMEM*E", "*E", "-*", "*--", "--", "-M", "E*E-", "MM*") + + def test_adjacent_layers_agree_on_their_shared_boundary(self): + for pattern, dp, tp, a2a in itertools.product( + self.PATTERNS, (1, 2), (1, 2), (False, True) + ): + with self.subTest(pattern=pattern, dp=dp, tp=tp, a2a=a2a): + layers = stages(pattern, dp=dp, tp=tp, a2a=a2a) + rows = Layout.sharded_over( + TokenAxis.ATTN_DP, axis_sizes=sizes(dp=dp, tp=tp) + ) + for before, after in zip(layers, layers[1:]): + out_of, into = before.edges[1], after.edges[0] + self.assertEqual(out_of.need.layout, rows) + self.assertEqual(out_of.residual_to, rows) + self.assertEqual(into.residual, rows) + self.assertEqual(into.produced.layout, rows) + produced = out_of.produced + may_leave = produced.always_leaves or produced.leaves_for_next_layer + self.assertEqual( + into.produced.group, produced.group if may_leave else None + ) + self.assertEqual( + into.produced.always_leaves, produced.always_leaves + ) + self.assertEqual( + into.produced.leaves_for_next_layer, + produced.leaves_for_next_layer, + ) + self.assertEqual(layers[0].edges[0].produced, StageOutput(rows)) + self.assertTrue(layers[0].enters_stack) + self.assertFalse(any(layer.enters_stack for layer in layers[1:])) + last = layers[-1].edges[1].produced + self.assertFalse(last.always_leaves or last.leaves_for_next_layer) + + def test_what_each_kind_of_boundary_carries(self): + # (pattern, boundary after layer 0): group, always_leaves, leaves_for_next_layer + cases = { + "M-": (SumGroup.ATTN_TP, True, False), + "*E": (SumGroup.ATTN_TP, True, False), + "MM": (SumGroup.ATTN_TP, False, True), + "-M": (SumGroup.TP, False, True), + "EM": (SumGroup.MOE_OUTPUT, False, True), + "--": (None, False, False), + "E-": (None, False, False), + } + for pattern, expected in cases.items(): + with self.subTest(pattern=pattern): + into = stages(pattern, tp=2)[1].edges[0].produced + self.assertEqual( + (into.group, into.always_leaves, into.leaves_for_next_layer), + expected, + ) + # Without attention TP a mixer's output is complete. + into = stages("M-", tp=1)[1].edges[0].produced + self.assertEqual(into, StageOutput(into.layout)) + # A MoE on this rank's own rows hands on a complete output; an a2a + # backend dispatches only the MoE, so an MLP still sums over TP. + into = stages("EM", tp=2, a2a=True)[1].edges[0].produced + self.assertEqual(into, StageOutput(into.layout)) + into = stages("-M", tp=2, a2a=True)[1].edges[0].produced + self.assertEqual( + (into.group, into.always_leaves, into.leaves_for_next_layer), + (SumGroup.TP, False, True), + ) + + def test_the_residual_follows_the_input_onto_a_finer_slice(self): + axis_sizes = sizes(dp=2, tp=2) + attention = Layout.sharded_over(TokenAxis.ATTN_DP, axis_sizes=axis_sizes) + local = Layout.sharded_over( + TokenAxis.ATTN_DP, TokenAxis.ATTN_TP_SCATTER, axis_sizes=axis_sizes + ) + full = Layout.sharded_over(axis_sizes=axis_sizes) + for need, during in ((local, local), (full, attention)): + stage = StageDecl(StageInput(need), StageOutput(need)) + into, out_of = stage_edges(previous=None, stage=stage, rows=attention) + self.assertEqual((into.residual, into.residual_to), (attention, during)) + self.assertEqual((out_of.residual, out_of.residual_to), (during, attention)) + + +class TestMixerExit(CustomTestCase): + """A mixer skips its output all-reduce when its output always leaves the sum + (to an FFN stage), and when it may leave it and the fused kernel takes it; + what it hands on says which.""" + + def test_decision_table(self): + tp_group = object() + attention = Layout.sharded_over(TokenAxis.ATTN_DP, axis_sizes=sizes(tp=2)) + for always, may, fuses in itertools.product((False, True), repeat=3): + if always and may: + continue + with self.subTest( + always_leaves=always, leaves_for_next_layer=may, fuses=fuses + ): + produced = StageOutput( + attention, + group=SumGroup.ATTN_TP if always or may else None, + always_leaves=always, + leaves_for_next_layer=may, + ) + communicator = SimpleNamespace( + _batch_steps=lambda batch: SimpleNamespace(ffn_output=produced), + should_fuse_mlp_allreduce_with_next_layer=MagicMock( + return_value=fuses + ), + ) + hidden = torch.ones(2, 4) + with get_parallel().override(tp_group=tp_group): + with MixerExit(communicator, None) as mixer_exit: + skipped = should_skip_mlp_all_reduce() + output = mixer_exit.finish(hidden) + self.assertFalse(should_skip_mlp_all_reduce()) + hands_on = may and fuses + self.assertEqual(mixer_exit.skips_reduction, always or hands_on) + self.assertEqual(skipped, always or hands_on) + if hands_on: + self.assertIsInstance(output, UnreducedOutput) + self.assertIs(output.group, tp_group) + else: + self.assertIs(output, hidden) + + +if __name__ == "__main__": + unittest.main()