diff --git a/python/sglang/srt/arg_groups/layernorm_sp_hook.py b/python/sglang/srt/arg_groups/layernorm_sp_hook.py new file mode 100644 index 000000000..d60194826 --- /dev/null +++ b/python/sglang/srt/arg_groups/layernorm_sp_hook.py @@ -0,0 +1,63 @@ +from __future__ import annotations + +import logging +from typing import TYPE_CHECKING, Optional + +if TYPE_CHECKING: + from sglang.srt.server_args import ServerArgs + +from sglang.srt.arg_groups.overrides import model_config_of, resolving_view + +logger = logging.getLogger(__name__) + + +def handle_layernorm_sp(server_args: ServerArgs) -> None: + """Validate --enable-layernorm-sp against the resolved parallelism config. + + Runs in the resolution pipeline rather than in the layers so a model that + never builds a LayerCommunicator rejects the flag instead of ignoring it. + """ + cfg = resolving_view(server_args) + if not cfg.enable_layernorm_sp: + return + architectures = model_config_of(server_args).hf_config.architectures + validate_layernorm_sp( + architecture=architectures[0] if architectures else None, + tp_size=cfg.tp_size, + enable_dp_attention=cfg.enable_dp_attention, + speculative_algorithm=cfg.speculative_algorithm, + ) + + +def validate_layernorm_sp( + *, + architecture: Optional[str], + tp_size: int, + enable_dp_attention: bool, + speculative_algorithm: Optional[str], +) -> None: + """Fail loud for unsupported / incompatible configs. Callers gate on the flag.""" + from sglang.srt.layers.layernorm_sp import SP_SUPPORTED_ARCHITECTURES + + if architecture not in SP_SUPPORTED_ARCHITECTURES: + raise ValueError( + "--enable-layernorm-sp is only supported for " + f"{sorted(SP_SUPPORTED_ARCHITECTURES)}; got {architecture}." + ) + if tp_size <= 1: + raise ValueError( + "--enable-layernorm-sp requires tp_size > 1: there is no sequence to " + "shard across a single TP rank." + ) + if enable_dp_attention: + raise ValueError( + "--enable-layernorm-sp is not compatible with --enable-dp-attention: " + "SP shards the sequence across the full TP group, which under DP " + "attention spans data-parallel groups holding different sequences." + ) + if speculative_algorithm is not None: + raise ValueError( + "--enable-layernorm-sp is not compatible with speculative decoding " + "(EAGLE/EAGLE3): the captured aux hidden states would be " + "sequence-sharded." + ) diff --git a/python/sglang/srt/arg_groups/pipeline.py b/python/sglang/srt/arg_groups/pipeline.py index cfe55dfca..5b9fd2364 100644 --- a/python/sglang/srt/arg_groups/pipeline.py +++ b/python/sglang/srt/arg_groups/pipeline.py @@ -318,6 +318,11 @@ def run_resolution_pipeline(server_args: Any) -> None: handle_speculative_decoding(server_args) + # After the speculative hook so speculative_algorithm is final. + from sglang.srt.arg_groups.layernorm_sp_hook import handle_layernorm_sp + + handle_layernorm_sp(server_args) + # Validate the CuteDSL A2A token budget now that num_tokens_per_req is final. validate_cutedsl_a2a_token_budget(server_args) diff --git a/python/sglang/srt/distributed/bootstrap.py b/python/sglang/srt/distributed/bootstrap.py index 570d6ba98..a4ab8318f 100644 --- a/python/sglang/srt/distributed/bootstrap.py +++ b/python/sglang/srt/distributed/bootstrap.py @@ -27,6 +27,7 @@ from sglang.srt.distributed.parallel_state import ( from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.environ import envs from sglang.srt.layers.dp_attention import initialize_dp_attention +from sglang.srt.layers.layernorm_sp import initialize_layernorm_sp from sglang.srt.platforms import current_platform from sglang.srt.runtime_context import ( get_disagg, @@ -280,6 +281,10 @@ def _init_parallel_groups( server_args=server_args, model_config=model_config, ) + initialize_layernorm_sp( + server_args=server_args, + model_config=model_config, + ) if is_npu(): register_sgl_tp_rank(gpu_id) diff --git a/python/sglang/srt/layers/communicator.py b/python/sglang/srt/layers/communicator.py index cbf100643..c5c35016a 100644 --- a/python/sglang/srt/layers/communicator.py +++ b/python/sglang/srt/layers/communicator.py @@ -31,6 +31,7 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) from sglang.srt.environ import envs +from sglang.srt.layers import layernorm_sp from sglang.srt.layers.attention.dsa.utils import ( dsa_use_prefill_cp, is_dsa_enable_prefill_cp, @@ -72,7 +73,7 @@ from sglang.srt.model_executor.cuda_graph_config import ( Phase, check_cuda_graph_backend, ) -from sglang.srt.model_executor.forward_batch_info import ForwardBatch +from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.runtime_context import ( get_exec, get_forward, @@ -483,6 +484,7 @@ class LayerCommunicator: force_layernorm_before_dp_gather: bool = False, enable_fused_ar_quant: bool = False, fused_ar_quant_keep_bf16: bool = False, + _is_sp_variant: bool = False, ): self.layer_scatter_modes = layer_scatter_modes self.input_layernorm = input_layernorm @@ -503,6 +505,30 @@ class LayerCommunicator: get_spec().speculative_algorithm ) + # Under LayerNorm SP the norm/residual run on the sequence shard with no + # collectives, so delegate to an all-SCATTERED sibling while the region is + # active. _is_sp_variant stops the sibling from building its own. + self._sp_variant: Optional[LayerCommunicator] = None + if not _is_sp_variant and layernorm_sp.layernorm_sp_enabled(): + self._sp_variant = LayerCommunicator( + layer_scatter_modes=LayerScatterModes( + layer_input_mode=ScatterMode.SCATTERED, + attn_mode=ScatterMode.SCATTERED, + mlp_mode=ScatterMode.SCATTERED, + middle_residual_mode=ScatterMode.SCATTERED, + layer_output_mode=ScatterMode.SCATTERED, + ), + input_layernorm=input_layernorm, + post_attention_layernorm=post_attention_layernorm, + allow_reduce_scatter=allow_reduce_scatter, + is_last_layer=is_last_layer, + qkv_latent_func=qkv_latent_func, + force_layernorm_before_dp_gather=force_layernorm_before_dp_gather, + enable_fused_ar_quant=enable_fused_ar_quant, + fused_ar_quant_keep_bf16=fused_ar_quant_keep_bf16, + _is_sp_variant=True, + ) + def _post_init_communicate(self): self._communicate_simple_fn = CommunicateSimpleFn.get_fn( input_mode=self.layer_scatter_modes.layer_input_mode, @@ -592,6 +618,24 @@ class LayerCommunicator: quant_format: str = "", post_residual_addition: Optional[torch.Tensor] = None, ): + # residual is None marks the first decoder layer, where the SP region + # opens: re-evaluated per forward so a crash mid-loop cannot leak into + # the next one. + if self._sp_variant is not None: + if residual is None: + get_forward().set( + "sp_active", forward_batch.forward_mode == ForwardMode.EXTEND + ) + if get_forward().sp_active: + hidden_states = layernorm_sp.sp_entry_scatter(hidden_states) + if get_forward().sp_active: + return self._sp_variant.prepare_attn( + hidden_states, + residual, + forward_batch, + quant_format, + post_residual_addition, + ) if get_attn_tp_context().input_scattered: hidden_states, residual = self._tp_reduce_scatter( hidden_states, @@ -788,6 +832,10 @@ class LayerCommunicator: forward_batch: ForwardBatch, cache=None, ): + if self._sp_variant is not None and get_forward().sp_active: + return self._sp_variant.prepare_mlp( + hidden_states, residual, forward_batch, cache + ) if cache is not None: self._context.cache = cache @@ -805,6 +853,10 @@ class LayerCommunicator: residual: torch.Tensor, forward_batch: ForwardBatch, ): + if self._sp_variant is not None and get_forward().sp_active: + return self._sp_variant.postprocess_layer( + hidden_states, residual, forward_batch + ) return self._communicate_summable_tensor_pair_fn( hidden_states=hidden_states, residual=residual, @@ -1028,6 +1080,22 @@ class CommunicateWithAllReduceAndLayerNormFn: ): return CommunicateWithAllReduceAndLayerNormFn._simple + if ( + hidden_states_input_mode == ScatterMode.SCATTERED + and residual_input_mode == ScatterMode.SCATTERED + and hidden_states_output_mode == ScatterMode.SCATTERED + and residual_output_mode == ScatterMode.SCATTERED + ): + # Megatron LayerNorm sequence parallelism (layers/layernorm_sp.py): + # activations stay sequence-sharded across the attn->mlp boundary, so + # there is nothing to gather or scatter here -- just the residual add + # plus LayerNorm on the local shard. The row-parallel o_proj already + # issued the reduce-scatter (g-bar) that the all-reduce would have + # done, and the g all-gather is fused into the next column-parallel + # linear. Distinct from the branch above because under pure TP + # attn_tp_size == tp_size > 1, so that gate does not fire. + return CommunicateWithAllReduceAndLayerNormFn._simple + if ( (hidden_states_input_mode == ScatterMode.TP_ATTN_FULL) and ( diff --git a/python/sglang/srt/layers/layernorm_sp.py b/python/sglang/srt/layers/layernorm_sp.py new file mode 100644 index 000000000..04e9e5722 --- /dev/null +++ b/python/sglang/srt/layers/layernorm_sp.py @@ -0,0 +1,245 @@ +# Copyright 2023-2024 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Megatron-style LayerNorm sequence parallelism (SP, arXiv:2205.05198). + +Under pure tensor parallelism the row-parallel ``all_reduce`` is algebraically a +``reduce_scatter`` (g-bar) followed by an ``all_gather`` (g). Splitting it that +way lets the LayerNorm / residual regions run on sequence-sharded activations -- +each rank holds 1/tp of the tokens -- which cuts the transient activation memory +of long-context prefill with no extra communication volume (all_reduce and +reduce_scatter+all_gather move the same bytes). + +Everything SP lives here so the feature stays decoupled from model code and from +``dp_attention.py``: + + - which models opt in (the Qwen3-dense allowlist) and config validation, + - the per-forward ``sp_active`` flag (a ForwardFlags bool) read at depth by the + participant linears and the ``LayerCommunicator``, + - the entry-scatter / exit-gather collectives, and + - the fused matmul + collective fast-paths for the participant linears. + +SP runs for prefill (EXTEND) only and is off by default; with the flag off, +nothing in this module executes. +""" + +from __future__ import annotations + +from typing import Optional + +import torch + +from sglang.srt.distributed import get_tp_group +from sglang.srt.runtime_context import get_flags, get_forward +from sglang.srt.utils.common import ceil_align + +# Architectures whose decoder layers route attention/MLP through +# ``LayerCommunicator`` with the standard participant linears, and for which SP +# has been validated. Other models reject --enable-layernorm-sp at construction. +# The mechanism is generic; extend the allowlist as families are validated. +SP_SUPPORTED_ARCHITECTURES = frozenset({"Qwen3ForCausalLM"}) + + +def initialize_layernorm_sp(*, server_args, model_config) -> None: + """Materialize ``flags.sp.enabled``; runs once per worker after distributed + setup, alongside ``initialize_dp_attention``.""" + architectures = model_config.hf_config.architectures + get_flags().sp.enabled = bool( + server_args.enable_layernorm_sp + and architectures + and architectures[0] in SP_SUPPORTED_ARCHITECTURES + ) + + +def layernorm_sp_enabled() -> bool: + return get_flags().sp.enabled + + +def runs_sp(forward_mode) -> bool: + """Whether this forward runs SP: an enabled model, prefill only. + + Code outside the CUDA-graph-captured region must use this and not + ``get_forward().sp_active``: Python writes made inside that region do not + re-execute on graph replay, so the flag is stale there. + """ + from sglang.srt.model_executor.forward_batch_info import ForwardMode + + return layernorm_sp_enabled() and forward_mode == ForwardMode.EXTEND + + +class _SPForwardState: + """Real (unpadded) token count of the current SP forward. + + An instance attribute, not a ForwardFlags int slot: this is read inside + torch.compile-traced linear code, where dynamo gives an attribute-source int + automatic-dynamic, while a dict-slot int recompiles per sequence length. + """ + + num_tokens: int = 0 + + +_sp_state = _SPForwardState() + + +def set_sp_num_tokens(num_tokens: int) -> None: + _sp_state.num_tokens = num_tokens + + +def sp_num_tokens() -> int: + return _sp_state.num_tokens + + +# --- entry scatter / exit gather (once per forward, at the boundary) ---------- +def sp_entry_scatter(hidden_states: torch.Tensor) -> torch.Tensor: + """Shard the replicated ``[M, h]`` hidden states along the token dim. + + Pads M up to a multiple of tp_size; the padding rows are dropped by the exit + gather. The input is replicated across the TP group, so this is a local slice. + """ + num_tokens = hidden_states.shape[0] + set_sp_num_tokens(num_tokens) + tp_group = get_tp_group() + tp_size = tp_group.world_size + if tp_size == 1: + return hidden_states + padded = ceil_align(num_tokens, tp_size) + if padded != num_tokens: + hidden_states = torch.nn.functional.pad( + hidden_states, (0, 0, 0, padded - num_tokens) + ) + return hidden_states.tensor_split(tp_size)[tp_group.rank_in_group].contiguous() + + +def sp_exit_gather(hidden_states: torch.Tensor, num_tokens: int) -> torch.Tensor: + """g: all-gather the per-rank shards back to the full sequence along dim 0, + then narrow to ``num_tokens`` (dropping the entry-scatter padding).""" + tp_group = get_tp_group() + tp_size = tp_group.world_size + if tp_size == 1: + return hidden_states[:num_tokens] + hidden_states = hidden_states.contiguous() + output = hidden_states.new_empty( + (hidden_states.shape[0] * tp_size, *hidden_states.shape[1:]) + ) + tp_group.all_gather_into_tensor(output, hidden_states) + return output[:num_tokens] + + +def maybe_exit_gather( + *, + hidden_states: torch.Tensor, + hidden_states_before_norm: Optional[torch.Tensor], + input_ids: Optional[torch.Tensor], + forward_mode, +) -> tuple[torch.Tensor, Optional[torch.Tensor]]: + """Undo the sequence sharding before the LM head, and leave the region. + + No-op unless this forward runs SP. The token count comes from ``input_ids`` + and the predicate from ``runs_sp``, so this stays correct on CUDA graph + replay, where the writes made inside the captured region do not re-execute. + """ + if not runs_sp(forward_mode) or input_ids is None: + return hidden_states, hidden_states_before_norm + num_tokens = input_ids.shape[0] + hidden_states = sp_exit_gather(hidden_states, num_tokens=num_tokens) + if hidden_states_before_norm is not None: + hidden_states_before_norm = sp_exit_gather( + hidden_states_before_norm, num_tokens=num_tokens + ) + get_forward().set("sp_active", False) + return hidden_states, hidden_states_before_norm + + +# --- fused matmul + collective fast-paths for the participant linears --------- +# Fused matmul+reduce-scatter (g-bar) and all-gather+matmul (g) overlap the +# collective with the GEMM. Availability is probed once at import; TP groups are +# registered for symmetric memory lazily by the fused ops on first use (the old +# enable_symm_mem_for_group is a deprecated no-op), so we only import the module +# to register the torch.ops.symm_mem namespace the probe checks. NVLink/NVSwitch. +try: + import torch.distributed._symmetric_memory # noqa: F401 + + _HAS_TORCH_SYMM_MEM_FUSED = hasattr( + torch.ops.symm_mem, "fused_matmul_reduce_scatter" + ) and hasattr(torch.ops.symm_mem, "fused_all_gather_matmul") +except Exception: + _HAS_TORCH_SYMM_MEM_FUSED = False + + +def sp_fused_matmul_eligible(linear) -> bool: + """Whether the torch symm_mem fused matmul+collective fast-path applies: the + ops are available and ``linear`` is unquantized, bias-free, bf16/fp16 (the + case the fused ops support). Depends only on static layer properties, so the + decision is identical across TP ranks. + """ + if not _HAS_TORCH_SYMM_MEM_FUSED or linear.bias is not None: + return False + from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod + + if not isinstance(linear.quant_method, UnquantizedLinearMethod): + return False + return linear.weight.dtype in (torch.bfloat16, torch.float16) + + +def column_parallel_g_matmul( + linear, input_parallel: torch.Tensor, bias +) -> torch.Tensor: + """Megatron SP g for a ColumnParallelLinear participant (qkv / gate_up). + + The input is this rank's sequence shard ``[M_pad/tp, K]``; all-gather it back + to the full sequence, matmul, and narrow to the real token count (recorded at + the entry scatter). Uses the fused symm-mem kernel when eligible (all-gather + + GEMM in one shot), else a plain all-gather + matmul. + """ + num_tokens = sp_num_tokens() + if sp_fused_matmul_eligible(linear): + group_name = get_tp_group().device_group.group_name + _, mm_outputs = torch.ops.symm_mem.fused_all_gather_matmul( + input_parallel.contiguous(), + [linear.weight.t()], + gather_dim=0, + group_name=group_name, + ) + return mm_outputs[0][:num_tokens] + gathered = sp_exit_gather(input_parallel, num_tokens=num_tokens) + return linear.quant_method.apply(linear, gathered, bias) + + +def row_parallel_gbar_matmul(linear, input_: torch.Tensor, bias) -> torch.Tensor: + """Megatron SP g-bar for a RowParallelLinear participant (o_proj / down). + + Computes ``input_ @ weight.T`` and reduce-scatters (sum) the result across the + TP group along dim 0, leaving this rank's ``[M_pad/tp, h]`` shard. The token + dim is padded to a multiple of tp_size (padding rows are zeros, dropped by the + exit gather). Uses the fused symm-mem kernel when eligible, else matmul + a + plain reduce-scatter. + """ + tp_size = linear.tp_size + x = input_.contiguous() + num_tokens = x.shape[0] + padded = ceil_align(num_tokens, tp_size) + if padded != num_tokens: + x = torch.nn.functional.pad(x, (0, 0, 0, padded - num_tokens)) + if sp_fused_matmul_eligible(linear): + group_name = get_tp_group().device_group.group_name + return torch.ops.symm_mem.fused_matmul_reduce_scatter( + x, + linear.weight.t(), + "sum", + scatter_dim=0, + group_name=group_name, + ) + full = linear.quant_method.apply(linear, x, bias) + output = full.new_empty((padded // tp_size, *full.shape[1:])) + get_tp_group().reduce_scatter_tensor(output, full) + return output diff --git a/python/sglang/srt/layers/linear.py b/python/sglang/srt/layers/linear.py index c900f533f..d84019fbd 100644 --- a/python/sglang/srt/layers/linear.py +++ b/python/sglang/srt/layers/linear.py @@ -25,6 +25,7 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import ( use_symmetric_memory, ) from sglang.srt.environ import envs +from sglang.srt.layers import layernorm_sp from sglang.srt.layers.dp_attention import ( is_allocation_symmetric, ) @@ -39,7 +40,7 @@ from sglang.srt.layers.parameter import ( _ColumnvLLMParameter, ) from sglang.srt.layers.utils import pad_or_narrow_weight -from sglang.srt.runtime_context import get_exec, get_parallel +from sglang.srt.runtime_context import get_exec, get_forward, get_parallel from sglang.srt.utils import get_bool_env_var, is_cpu, is_hip, is_npu, set_weight_attrs if TYPE_CHECKING: @@ -492,6 +493,14 @@ class ColumnParallelLinear(LinearBase): def forward(self, input_): bias = self.bias if not self.skip_bias_add else None + # Megatron SP "g": the input is this rank's [M_pad/tp, K] sequence shard; + # all-gather to the full sequence and matmul. Participants (qkv/gate_up) + # have gather_output=False, so there is no output all-gather to reconcile. + if get_forward().sp_active and self.tp_size > 1: + output = layernorm_sp.column_parallel_g_matmul(self, input_, bias) + output_bias = self.bias if self.skip_bias_add else None + return output, output_bias + # Matrix multiply. assert self.quant_method is not None output_parallel = self.quant_method.apply(self, input_, bias) @@ -1621,6 +1630,21 @@ class RowParallelLinear(LinearBase): # Only fuse bias add into GEMM for rank 0 (this ensures that # bias will not get added more than once in TP>1 case) bias_ = None if (self.tp_rank > 0 or self.skip_bias_add) else self.bias + + # Megatron SP "g-bar": reduce-scatter along the token dim instead of + # all-reduce, leaving the output sharded for the next SP LayerNorm region. + # Fires regardless of reduce_results: o_proj / down are built + # reduce_results=False, so under SP the linear owns the reduction. + if ( + get_forward().sp_active + and self.tp_size > 1 + and not skip_all_reduce + and output_tensor is None + ): + output = layernorm_sp.row_parallel_gbar_matmul(self, input_parallel, bias_) + output_bias = self.bias if self.skip_bias_add else None + return output, output_bias + if self.use_dp_attention_reduce: symm_ctx = use_symmetric_memory(get_parallel().attn_tp_group) else: diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py index 8de83128a..03d09c9ca 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -28,6 +28,7 @@ from sglang.srt.beam_search.logits_capture import BeamLogitsCapture from sglang.srt.distributed import get_tp_group from sglang.srt.distributed.device_communicators import triton_symm_mem_ag from sglang.srt.environ import envs +from sglang.srt.layers import layernorm_sp from sglang.srt.layers.aux_hidden_states import ( AuxHiddenStates, pack_aux_hidden_states, @@ -439,6 +440,15 @@ class LogitsProcessor(nn.Module): if _autotune_run_lm_head is False: return LogitsProcessorOutput(next_token_logits=None) + # Under LayerNorm SP the decoder loop leaves these sequence-sharded; undo + # that before the LM head, which must not participate. + hidden_states, hidden_states_before_norm = layernorm_sp.maybe_exit_gather( + hidden_states=hidden_states, + hidden_states_before_norm=hidden_states_before_norm, + input_ids=input_ids, + forward_mode=logits_metadata.forward_mode, + ) + # Multi-item scoring only for prefill-only requests with pre-computed indices. if multi_item_delimiter_indices is not None and logits_metadata.is_prefill_only: return self.compute_logprobs_for_multi_item_scoring( diff --git a/python/sglang/srt/runtime_context.py b/python/sglang/srt/runtime_context.py index f379eef83..c9cd6e120 100644 --- a/python/sglang/srt/runtime_context.py +++ b/python/sglang/srt/runtime_context.py @@ -522,6 +522,15 @@ class DpFlags(_FlagGroupBase): buffer_device: Any = None +@dataclasses.dataclass +class SpFlags(_FlagGroupBase): + """LayerNorm sequence-parallelism flags, materialized by + ``initialize_layernorm_sp`` (after distributed setup; reads the model + config). See ``layers.layernorm_sp``.""" + + enabled: bool = False + + @dataclasses.dataclass class Flags(_FlagGroupBase): """Root of the runtime-flags tier. @@ -529,12 +538,13 @@ class Flags(_FlagGroupBase): Resolved configuration lives in the config bags below (projected from the declarations at publish) — this tier only carries genuine runtime state whose value is not a function of the configuration alone, grouped - by lifecycle (``capture``) or subsystem (``moe`` / ``dp``). + by lifecycle (``capture``) or subsystem (``moe`` / ``dp`` / ``sp``). """ capture: CaptureFlags = dataclasses.field(default_factory=CaptureFlags) moe: MoeFlags = dataclasses.field(default_factory=MoeFlags) dp: DpFlags = dataclasses.field(default_factory=DpFlags) + sp: SpFlags = dataclasses.field(default_factory=SpFlags) @dataclasses.dataclass @@ -606,6 +616,8 @@ class ForwardFlags: "fuse_mlp_allreduce": False, "mlp_reduce_scatter": False, "flashinfer_trtllm_bypass": False, + # LayerNorm sequence parallelism region; see layers/layernorm_sp.py. + "sp_active": False, } # Read/written inside compiled graphs (vocab embedding, communicator, @@ -620,6 +632,7 @@ class ForwardFlags: "fuse_mlp_allreduce", "mlp_reduce_scatter", "flashinfer_trtllm_bypass", + "sp_active", } ) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index fb2661b57..49ba9ee32 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -1124,6 +1124,15 @@ class ServerArgs: "Shard dense MLP weights across the attention TP group under DP attention.", NS("parallel"), ] = False + enable_layernorm_sp: A[ + bool, + "Enable Megatron-style sequence parallelism (arXiv:2205.05198) for the " + "LayerNorm/residual regions under pure tensor parallelism: the row-parallel " + "all-reduce becomes reduce-scatter + all-gather, so LayerNorm runs on " + "sequence-sharded activations with no extra communication volume. " + "Prefill only; Qwen3 dense; requires tp_size > 1 and NVLink/NVSwitch.", + NS("parallel"), + ] = False disable_attn_tp_gather: A[ bool, "Disable scheduler-side attn_tp_gather (the upstream SP path " diff --git a/test/registered/models_e2e/test_layernorm_sp.py b/test/registered/models_e2e/test_layernorm_sp.py new file mode 100644 index 000000000..7103ae63b --- /dev/null +++ b/test/registered/models_e2e/test_layernorm_sp.py @@ -0,0 +1,53 @@ +"""E2E accuracy test for --enable-layernorm-sp (Megatron LayerNorm sequence +parallelism). + +Launches a Qwen3 dense model with tp=2 and the SP flag on, then checks GSM8K +accuracy stays healthy -- SP re-associates the row-parallel all-reduce as +reduce-scatter + all-gather and runs the norm/residual regions on sequence +shards, so a correct implementation matches the non-SP result within +floating-point-reordering noise. tp>1 is required for SP to engage. +""" + +import unittest + +from sglang.srt.utils import kill_process_tree +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kits.eval_accuracy_kit import GSM8KMixin +from sglang.test.test_utils import ( + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + DEFAULT_URL_FOR_TEST, + CustomTestCase, + popen_launch_server, +) + +register_cuda_ci(est_time=420, stage="base-b", runner_config="2-gpu-large") + +# Qwen3 dense (architecture "Qwen3ForCausalLM") is the SP allowlist entry. +LAYERNORM_SP_MODEL = "Qwen/Qwen3-8B" + + +class TestLayerNormSPAccuracy(CustomTestCase, GSM8KMixin): + gsm8k_score_threshold = 0.85 + gsm8k_num_examples = 200 + + @classmethod + def setUpClass(cls): + cls.model = LAYERNORM_SP_MODEL + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=["--tp", "2", "--enable-layernorm-sp"], + ) + + @classmethod + def tearDownClass(cls): + if hasattr(cls, "process") and cls.process: + kill_process_tree(cls.process.pid) + + # test_gsm8k is provided by GSM8KMixin (asserts score >= gsm8k_score_threshold). + + +if __name__ == "__main__": + unittest.main(verbosity=3) diff --git a/test/registered/unit/layers/test_layernorm_sp.py b/test/registered/unit/layers/test_layernorm_sp.py new file mode 100644 index 000000000..23925d51a --- /dev/null +++ b/test/registered/unit/layers/test_layernorm_sp.py @@ -0,0 +1,125 @@ +"""Unit tests for srt/layers/layernorm_sp (Megatron LayerNorm sequence parallelism). + +Covers the pure logic that gates SP -- the Qwen3 allowlist, the config guards, and +the prefill-only activation rule -- without launching a server. The collectives and +fused matmul fast-paths need a real TP group and are covered by the e2e test. +""" + +import unittest +from types import SimpleNamespace + +from sglang.srt.arg_groups.layernorm_sp_hook import validate_layernorm_sp +from sglang.srt.layers import layernorm_sp +from sglang.srt.model_executor.forward_batch_info import ForwardMode +from sglang.srt.runtime_context import get_flags, get_forward, reset_context +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=10, suite="base-a-test-cpu") + + +def _initialize(*, enable=True, arch="Qwen3ForCausalLM"): + layernorm_sp.initialize_layernorm_sp( + server_args=SimpleNamespace(enable_layernorm_sp=enable), + model_config=SimpleNamespace( + hf_config=SimpleNamespace(architectures=[arch] if arch else []) + ), + ) + + +class TestLayerNormSPGating(CustomTestCase): + def tearDown(self): + reset_context() + + def test_initialize_enables_only_for_allowlisted_arch(self): + _initialize(enable=True, arch="Qwen3ForCausalLM") + self.assertTrue(layernorm_sp.layernorm_sp_enabled()) + # Flag on but unsupported architecture -> stays off. + _initialize(enable=True, arch="LlamaForCausalLM") + self.assertFalse(layernorm_sp.layernorm_sp_enabled()) + # Supported architecture but flag off -> stays off. + _initialize(enable=False, arch="Qwen3ForCausalLM") + self.assertFalse(layernorm_sp.layernorm_sp_enabled()) + + def test_defaults_off_before_initialization(self): + # A process that never runs initialize_layernorm_sp must not enable SP. + self.assertFalse(layernorm_sp.layernorm_sp_enabled()) + + def test_runs_sp_only_on_extend(self): + with get_flags().sp.override(enabled=True): + self.assertTrue(layernorm_sp.runs_sp(ForwardMode.EXTEND)) + # SP is prefill-only; decode must never engage it. + self.assertFalse(layernorm_sp.runs_sp(ForwardMode.DECODE)) + # A disabled model never activates, even on EXTEND. + self.assertFalse(layernorm_sp.runs_sp(ForwardMode.EXTEND)) + + def test_runs_sp_ignores_the_active_flag(self): + """Regression: the exit gather must not key off ``sp_active``. + + ``sp_active`` is written by Python inside the CUDA-graph-captured region, + so it is stale on graph replay. Callers outside that region (the exit + gather in LogitsProcessor) recompute with ``runs_sp`` instead; if that ever + starts consulting the flag, a replayed prefill skips the gather and feeds + sequence-sharded hidden states to the LM head. + """ + with get_flags().sp.override(enabled=True): + get_forward().set("sp_active", False) + self.assertTrue(layernorm_sp.runs_sp(ForwardMode.EXTEND)) + reset_context() + + +class TestLayerNormSPValidation(CustomTestCase): + """``validate_layernorm_sp`` is pure; pass config in directly.""" + + VALID = dict( + architecture="Qwen3ForCausalLM", + tp_size=2, + enable_dp_attention=False, + speculative_algorithm=None, + ) + + def test_valid_config_passes(self): + validate_layernorm_sp(**self.VALID) # must not raise + + def test_rejects_unsupported_arch(self): + with self.assertRaisesRegex(ValueError, "only supported"): + validate_layernorm_sp(**{**self.VALID, "architecture": "LlamaForCausalLM"}) + + def test_rejects_tp_size_one(self): + with self.assertRaisesRegex(ValueError, "tp_size"): + validate_layernorm_sp(**{**self.VALID, "tp_size": 1}) + + def test_rejects_dp_attention(self): + with self.assertRaisesRegex(ValueError, "dp-attention"): + validate_layernorm_sp(**{**self.VALID, "enable_dp_attention": True}) + + def test_rejects_speculative(self): + with self.assertRaisesRegex(ValueError, "speculative"): + validate_layernorm_sp(**{**self.VALID, "speculative_algorithm": "EAGLE3"}) + + +class TestLayerNormSPActiveFlag(CustomTestCase): + def tearDown(self): + reset_context() + + def test_active_flag_is_registered_on_forward_flags(self): + """``sp_active`` must stay a registered ForwardFlags slot. + + ``set()`` rejects names missing from ``ForwardFlags._DEFAULTS``, so this + fails if the slot is dropped, and it must also be in ``_GRAPH_VISIBLE`` + because the participant linears read it under CUDA graph capture. + """ + from sglang.srt.runtime_context import ForwardFlags + + self.assertIn("sp_active", ForwardFlags._DEFAULTS) + self.assertIn("sp_active", ForwardFlags._GRAPH_VISIBLE) + + self.assertFalse(get_forward().sp_active) # default + get_forward().set("sp_active", True) + self.assertTrue(get_forward().sp_active) + get_forward().set("sp_active", False) + self.assertFalse(get_forward().sp_active) + + +if __name__ == "__main__": + unittest.main()