Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
44 changes: 43 additions & 1 deletion python/sglang/srt/layers/boundary_layout.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down
159 changes: 143 additions & 16 deletions python/sglang/srt/layers/communicator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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``."""
Expand Down Expand Up @@ -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`` /
Expand Down Expand Up @@ -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()
Expand All @@ -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:
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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,
),
Expand Down Expand Up @@ -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,
Expand Down
Loading
Loading