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()