From 2d0e94e3a36cbf125d52eaf20bfb26117a6e72ca Mon Sep 17 00:00:00 2001 From: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Date: Mon, 21 Sep 2026 12:22:59 -0700 Subject: [PATCH] Check the topology identities where the layout is written, and build at the published widths (#40340) --- benchmark/hf3fs/bench_zerocopy.py | 9 +- .../all_reduce/benchmark_all_reduce.py | 4 +- .../all_reduce/benchmark_fused_ar_rms_amd.py | 4 +- .../benchmark_fused_ar_rms_quant_amd.py | 4 +- .../kernels/all_reduce/benchmark_mscclpp.py | 4 +- .../all_reduce/benchmark_torch_symm_mem.py | 4 +- .../benchmark_fused_collective.py | 7 +- .../benchmark_sglang_fused_moe_triton.py | 7 +- ...nchmark_vllm_vs_sglang_fused_moe_triton.py | 7 +- .../usage/modelopt_quantize_and_export.py | 7 +- python/sglang/srt/arg_groups/arg_utils.py | 2 +- .../srt/disaggregation/encoder/server.py | 28 +- python/sglang/srt/distributed/bootstrap.py | 31 +- .../sglang/srt/distributed/parallel_state.py | 80 +-- python/sglang/srt/layers/dp_attention.py | 76 ++- python/sglang/srt/layers/sampler.py | 7 +- .../managers/scheduler_components/dp_attn.py | 4 +- .../srt/model_executor/forward_batch_info.py | 7 +- .../sglang/srt/models/inkling_common/moe.py | 6 + python/sglang/srt/models/kimi_k3.py | 3 +- python/sglang/srt/runtime_context.py | 162 ++++++- python/sglang/srt/weight_cache/daemon.py | 11 +- python/sglang/test/layer_ut_utils.py | 10 +- python/sglang/test/test_utils.py | 22 + test/manual/ep/test_flashinfer_dispatcher.py | 7 +- test/manual/test_custom_allreduce.py | 8 +- test/manual/test_quick_allreduce.py | 8 +- test/manual/test_weight_loader_v2_equiv.py | 4 +- .../cpu/test_rank_consensus_checker.py | 13 +- .../e2e/pp/test_pp_spec_embed_scan.py | 9 +- test/registered/eplb/test_lplb_distributed.py | 7 +- .../test_deterministic_reduce_scatter.py | 4 +- .../ops/moe/test_fused_swiglu_epilogue.py | 9 +- .../layers/mamba/test_mamba2_mixer.py | 4 +- .../ops/test_aiter_allreduce_fusion_amd.py | 5 +- .../unit/distributed/test_parallel_state.py | 16 +- .../attention/test_vision_strided_qkv.py | 4 +- .../test_dsa_layer_split_broadcast.py | 8 +- .../mem_cache/test_page_interleave_shard.py | 8 +- .../model_loader/test_presharded_loader.py | 4 +- .../unit/models/test_glm5_next_bfg_fusion.py | 13 +- .../test_shared_experts_fusion_gates.py | 10 +- test/registered/unit/test_runtime_context.py | 457 +++++++++++++++--- 43 files changed, 843 insertions(+), 261 deletions(-) diff --git a/benchmark/hf3fs/bench_zerocopy.py b/benchmark/hf3fs/bench_zerocopy.py index a4ed66912..8559e45b7 100644 --- a/benchmark/hf3fs/bench_zerocopy.py +++ b/benchmark/hf3fs/bench_zerocopy.py @@ -5,10 +5,10 @@ import torch from tqdm import tqdm from sglang.srt.distributed import ( - get_world_group, init_distributed_environment, initialize_model_parallel, ) +from sglang.srt.distributed.parallel_state import get_world_group from sglang.srt.managers.cache_controller import ( HiCacheController, PrefetchOperation, @@ -17,6 +17,7 @@ from sglang.srt.managers.cache_controller import ( from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator from sglang.srt.mem_cache.memory_pool import MHATokenToKVPool from sglang.srt.mem_cache.pool_host.mha import MHATokenToKVPoolHost +from sglang.test.test_utils import publish_build_topology init_distributed_environment( world_size=1, @@ -26,10 +27,8 @@ init_distributed_environment( backend="gloo", ) -initialize_model_parallel( - tensor_model_parallel_size=1, - pipeline_model_parallel_size=1, -) +publish_build_topology() +initialize_model_parallel() group = get_world_group().cpu_group diff --git a/benchmark/kernels/all_reduce/benchmark_all_reduce.py b/benchmark/kernels/all_reduce/benchmark_all_reduce.py index 30a7733f3..41dda1930 100644 --- a/benchmark/kernels/all_reduce/benchmark_all_reduce.py +++ b/benchmark/kernels/all_reduce/benchmark_all_reduce.py @@ -21,6 +21,7 @@ from sglang.srt.distributed.parallel_state import ( init_distributed_environment, initialize_model_parallel, ) +from sglang.test.test_utils import publish_build_topology def parse_args(): @@ -85,7 +86,8 @@ def init_dist(backend: str): distributed_init_method=distributed_init_method, local_rank=rank, ) - initialize_model_parallel(tensor_model_parallel_size=world_size) + publish_build_topology(world_rank=rank, tp_size=world_size) + initialize_model_parallel() return dist.group.WORLD diff --git a/benchmark/kernels/all_reduce/benchmark_fused_ar_rms_amd.py b/benchmark/kernels/all_reduce/benchmark_fused_ar_rms_amd.py index 1fa3819cc..107533fb1 100644 --- a/benchmark/kernels/all_reduce/benchmark_fused_ar_rms_amd.py +++ b/benchmark/kernels/all_reduce/benchmark_fused_ar_rms_amd.py @@ -41,6 +41,7 @@ from sglang.srt.distributed.parallel_state import ( initialize_model_parallel, set_custom_all_reduce, ) +from sglang.test.test_utils import publish_build_topology Shape = Tuple[int, int] @@ -381,7 +382,8 @@ def main(): distributed_init_method="env://", backend="nccl", ) - initialize_model_parallel(tensor_model_parallel_size=world_size) + publish_build_topology(world_rank=rank, tp_size=world_size) + initialize_model_parallel() prefill_shapes = parse_shapes(args.prefill_shapes) decode_shapes = parse_shapes(args.decode_shapes) diff --git a/benchmark/kernels/all_reduce/benchmark_fused_ar_rms_quant_amd.py b/benchmark/kernels/all_reduce/benchmark_fused_ar_rms_quant_amd.py index 85cabe94c..4c02a5df4 100644 --- a/benchmark/kernels/all_reduce/benchmark_fused_ar_rms_quant_amd.py +++ b/benchmark/kernels/all_reduce/benchmark_fused_ar_rms_quant_amd.py @@ -47,6 +47,7 @@ from sglang.srt.distributed.parallel_state import ( initialize_model_parallel, set_custom_all_reduce, ) +from sglang.test.test_utils import publish_build_topology Shape = Tuple[int, int] FP8_DTYPE = torch.float8_e4m3fnuz @@ -400,7 +401,8 @@ def main() -> None: distributed_init_method="env://", backend="nccl", ) - initialize_model_parallel(tensor_model_parallel_size=world_size) + publish_build_topology(world_rank=rank, tp_size=world_size) + initialize_model_parallel() if rank == 0: print( diff --git a/benchmark/kernels/all_reduce/benchmark_mscclpp.py b/benchmark/kernels/all_reduce/benchmark_mscclpp.py index 5b72d2f46..87dfd17e5 100644 --- a/benchmark/kernels/all_reduce/benchmark_mscclpp.py +++ b/benchmark/kernels/all_reduce/benchmark_mscclpp.py @@ -30,6 +30,7 @@ from sglang.srt.distributed.parallel_state import ( initialize_model_parallel, set_mscclpp_all_reduce, ) +from sglang.test.test_utils import publish_build_topology def torch_allreduce(torch_input: torch.Tensor, group: ProcessGroup) -> torch.Tensor: @@ -173,7 +174,8 @@ if __name__ == "__main__": rank=rank, local_rank=rank % 8, ) - initialize_model_parallel(tensor_model_parallel_size=world_size) + publish_build_topology(world_rank=rank, tp_size=world_size) + initialize_model_parallel() group = get_tensor_model_parallel_group().device_group cpu_group = get_tensor_model_parallel_group().cpu_group pynccl_comm = get_tensor_model_parallel_group().pynccl_comm diff --git a/benchmark/kernels/all_reduce/benchmark_torch_symm_mem.py b/benchmark/kernels/all_reduce/benchmark_torch_symm_mem.py index 5bdb7f5d6..3d2c65870 100644 --- a/benchmark/kernels/all_reduce/benchmark_torch_symm_mem.py +++ b/benchmark/kernels/all_reduce/benchmark_torch_symm_mem.py @@ -44,6 +44,7 @@ from sglang.srt.distributed.parallel_state import ( initialize_model_parallel, set_torch_symm_mem_all_reduce, ) +from sglang.test.test_utils import publish_build_topology from sglang.utils import is_in_ci IS_CI = is_in_ci() @@ -188,7 +189,8 @@ if __name__ == "__main__": rank=rank, local_rank=rank % 8, ) - initialize_model_parallel(tensor_model_parallel_size=world_size) + publish_build_topology(world_rank=rank, tp_size=world_size) + initialize_model_parallel() group = get_tensor_model_parallel_group().device_group cpu_group = get_tensor_model_parallel_group().cpu_group pynccl_comm = get_tensor_model_parallel_group().pynccl_comm diff --git a/benchmark/kernels/flashinfer_allreduce_fusion/benchmark_fused_collective.py b/benchmark/kernels/flashinfer_allreduce_fusion/benchmark_fused_collective.py index 2905a62a1..9af15ac87 100644 --- a/benchmark/kernels/flashinfer_allreduce_fusion/benchmark_fused_collective.py +++ b/benchmark/kernels/flashinfer_allreduce_fusion/benchmark_fused_collective.py @@ -30,14 +30,16 @@ import torch.distributed as dist # type: ignore from sglang.kernels.ops.quantization.fp8_kernel import fp8_dtype as SGLANG_FP8_DTYPE from sglang.kernels.ops.quantization.fp8_kernel import static_quant_fp8 -from sglang.srt.distributed import get_tp_group, tensor_model_parallel_all_reduce +from sglang.srt.distributed import tensor_model_parallel_all_reduce from sglang.srt.distributed.parallel_state import ( cleanup_dist_env_and_memory, + get_tp_group, graph_capture, init_distributed_environment, initialize_model_parallel, ) from sglang.srt.layers.layernorm import RMSNorm # noqa +from sglang.test.test_utils import publish_build_topology try: from sgl_kernel import fused_add_rmsnorm as SGL_FUSED_ADD_RMS_NORM @@ -1178,7 +1180,8 @@ def main(): local_rank=rank, backend="nccl", ) - initialize_model_parallel(tensor_model_parallel_size=world_size) + publish_build_topology(world_rank=rank, tp_size=world_size) + initialize_model_parallel() # Validate world size (must be > 1 for collective operations) if world_size <= 1: diff --git a/benchmark/kernels/fused_moe_triton/benchmark_sglang_fused_moe_triton.py b/benchmark/kernels/fused_moe_triton/benchmark_sglang_fused_moe_triton.py index 4515ff53b..6934650ad 100644 --- a/benchmark/kernels/fused_moe_triton/benchmark_sglang_fused_moe_triton.py +++ b/benchmark/kernels/fused_moe_triton/benchmark_sglang_fused_moe_triton.py @@ -26,6 +26,7 @@ from sglang.srt.layers.moe.topk import ( select_experts, ) from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler +from sglang.test.test_utils import publish_build_topology def fused_moe_triton_api( @@ -227,10 +228,8 @@ def main(): backend="nccl" if torch.cuda.is_available() else "gloo", ) - initialize_model_parallel( - tensor_model_parallel_size=1, - expert_model_parallel_size=1, - ) + publish_build_topology() + initialize_model_parallel() model_config = get_model_config(args.model, args.tp_size, args.ep_size) benchmark.run( diff --git a/benchmark/kernels/fused_moe_triton/benchmark_vllm_vs_sglang_fused_moe_triton.py b/benchmark/kernels/fused_moe_triton/benchmark_vllm_vs_sglang_fused_moe_triton.py index fc100ce50..1adecbeae 100644 --- a/benchmark/kernels/fused_moe_triton/benchmark_vllm_vs_sglang_fused_moe_triton.py +++ b/benchmark/kernels/fused_moe_triton/benchmark_vllm_vs_sglang_fused_moe_triton.py @@ -15,6 +15,7 @@ from sglang.srt.distributed.parallel_state import ( from sglang.srt.layers.moe.moe_runner.triton_utils.fused_moe import ( fused_moe as fused_moe_sglang, ) +from sglang.test.test_utils import publish_build_topology from .common_utils import get_model_config @@ -243,10 +244,8 @@ def main(): backend="nccl" if torch.cuda.is_available() else "gloo", ) - initialize_model_parallel( - tensor_model_parallel_size=1, - pipeline_model_parallel_size=1, - ) + publish_build_topology() + initialize_model_parallel() shape_configs = get_model_config(args.model, args.tp_size, args.ep_size) benchmark.run( diff --git a/examples/usage/modelopt_quantize_and_export.py b/examples/usage/modelopt_quantize_and_export.py index b3f26ede4..a8c2c8ad8 100755 --- a/examples/usage/modelopt_quantize_and_export.py +++ b/examples/usage/modelopt_quantize_and_export.py @@ -21,6 +21,7 @@ from sglang.srt.distributed.parallel_state import ( initialize_model_parallel, ) from sglang.srt.model_loader.loader import get_model_loader +from sglang.test.test_utils import publish_build_topology def _validate_export(export_dir: str) -> bool: @@ -113,10 +114,8 @@ def quantize_and_export_model( local_rank=0, backend="nccl" if device == "cuda" else "gloo", ) - initialize_model_parallel( - tensor_model_parallel_size=1, - pipeline_model_parallel_size=1, - ) + publish_build_topology() + initialize_model_parallel() # Configure model loading with ModelOpt quantization and export model_config = ModelConfig( diff --git a/python/sglang/srt/arg_groups/arg_utils.py b/python/sglang/srt/arg_groups/arg_utils.py index 0c54e930f..1fbf53b30 100644 --- a/python/sglang/srt/arg_groups/arg_utils.py +++ b/python/sglang/srt/arg_groups/arg_utils.py @@ -110,7 +110,7 @@ class Derived(msgspec.Struct, frozen=True): Every declaration carries ``fn`` today, the parallel quotients included: they are a function of the configured leaves, so they are computed at publish like the rest. What is special about them is not how they are - computed but that a stamp can move one afterwards -- an elastic scale-up + computed but that a stamp can move one afterwards -- ``initialize_dp_attention`` restamps ``attn_dp_size`` -- which ``ParallelContext`` answers above the published leaf. """ diff --git a/python/sglang/srt/disaggregation/encoder/server.py b/python/sglang/srt/disaggregation/encoder/server.py index b4b2aef1e..67f1abde0 100644 --- a/python/sglang/srt/disaggregation/encoder/server.py +++ b/python/sglang/srt/disaggregation/encoder/server.py @@ -589,10 +589,32 @@ class MMEncoder: distributed_init_method=dist_init_method, local_rank=rank, ) - initialize_model_parallel( - tensor_model_parallel_size=get_parallel().tp_size, - attention_context_model_parallel_size=get_parallel().attn_cp_size, + # The encoder serves the vision tower on a world of its own: `tp_size` + # ranks wide, with no pipeline, no expert or MoE-DP dimension and no + # decode context parallelism, whatever the generation side published. + # That has always been the layout it builds; stating it is what stops + # the context from answering with the other side's topology while these + # groups answer with this one. + parallel = get_parallel() + attn_cp_size = parallel.attn_cp_size + attn_tp_size = parallel.tp_size // attn_cp_size + attn_cp_rank, attn_tp_rank = divmod(rank, attn_tp_size) + parallel.override_permanently( + tp_rank=rank, + pp_size=1, + pp_rank=0, + attn_dp_size=1, + attn_dp_rank=0, + attn_tp_size=attn_tp_size, + attn_tp_rank=attn_tp_rank, + attn_cp_rank=attn_cp_rank, + attn_dcp_size=1, + moe_ep_size=1, + moe_ep_rank=0, + moe_dp_size=1, + moe_tp_size=parallel.tp_size, ) + initialize_model_parallel() initialize_dp_attention(server_args, self.model_config) self.model = load_model( diff --git a/python/sglang/srt/distributed/bootstrap.py b/python/sglang/srt/distributed/bootstrap.py index d45c48d39..6c91c1efc 100644 --- a/python/sglang/srt/distributed/bootstrap.py +++ b/python/sglang/srt/distributed/bootstrap.py @@ -106,15 +106,6 @@ def init_torch_distributed( server_args=server_args, model_config=model_config, gpu_id=ps.gpu_id, - tp_rank=ps.tp_rank, - tp_size=ps.tp_size, - pp_rank=ps.pp_rank, - pp_size=ps.pp_size, - attn_dp_size=ps.attn_dp_size, - attn_cp_size=ps.attn_cp_size, - moe_ep_size=ps.moe_ep_size, - moe_dp_size=ps.moe_dp_size, - dcp_size=ps.attn_dcp_size, ) # Pre-warm NCCL/RCCL/HCCL to eliminate cold-start latency in first request @@ -255,19 +246,13 @@ def _init_parallel_groups( server_args: ServerArgs, model_config: ModelConfig, gpu_id: int, - tp_rank: int, - tp_size: int, - pp_rank: int, - pp_size: int, - attn_dp_size: int, - attn_cp_size: int, - moe_ep_size: int, - moe_dp_size: int, - dcp_size: int, ) -> None: + parallel = get_parallel() + tp_size, pp_size = parallel.tp_size, parallel.pp_size + tp_rank, pp_rank = parallel.tp_rank, parallel.pp_rank is_ep_joiner = get_exec().moe.is_ep_joiner is_scale_joiner = get_exec().moe.is_ep_scale_joiner - rank_offset = get_parallel().ep_join_rank_offset if is_scale_joiner else 0 + rank_offset = parallel.ep_join_rank_offset if is_scale_joiner else 0 world_size = ( rank_offset + tp_size * pp_size if is_scale_joiner else tp_size * pp_size ) @@ -285,14 +270,6 @@ def _init_parallel_groups( max_world_size=get_parallel().max_ep_size, ) initialize_model_parallel( - tensor_model_parallel_size=tp_size, - attention_data_parallel_size=attn_dp_size, - pipeline_model_parallel_size=pp_size, - expert_model_parallel_size=moe_ep_size, - attention_context_model_parallel_size=attn_cp_size, - moe_data_model_parallel_size=moe_dp_size, - decode_context_parallel_size=dcp_size, - shared_experts_tensor_parallel_size=get_parallel().shared_experts_tp_size, duplicate_tp_group=get_disagg().enable_pdmux, enable_symm_mem=get_exec().comm.enable_symm_mem, # Only WORLD is extended during scale-up. The joiner's model-parallel diff --git a/python/sglang/srt/distributed/parallel_state.py b/python/sglang/srt/distributed/parallel_state.py index abf9024ca..4c02ebb21 100644 --- a/python/sglang/srt/distributed/parallel_state.py +++ b/python/sglang/srt/distributed/parallel_state.py @@ -50,6 +50,7 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo ) from sglang.srt.platforms.device_mixin import _DEVICE_TO_DISTRIBUTED_BACKEND from sglang.srt.runtime_context import ( + _validate_parallel, derive_parallel_widths, get_global_dwdp_manager, get_parallel, @@ -2515,44 +2516,41 @@ def init_distributed_environment( def initialize_model_parallel( - tensor_model_parallel_size: int = 1, - expert_model_parallel_size: int = 1, - pipeline_model_parallel_size: int = 1, - attention_data_parallel_size: int = 1, - attention_context_model_parallel_size: int = 1, - moe_data_model_parallel_size: int = 1, - decode_context_parallel_size: int = 1, backend: Optional[str] = None, duplicate_tp_group: bool = False, enable_symm_mem: bool = False, recovered_rank: bool = False, rank_offset: int = 0, max_world_size: Optional[int] = None, - shared_experts_tensor_parallel_size: Optional[int] = None, ) -> None: """ - Initialize model parallel groups. + Initialize model parallel groups at the published widths. - Arguments: - tensor_model_parallel_size: number of GPUs used for tensor model - parallelism. - expert_model_parallel_size: number of GPUs used for expert model - parallelism. - pipeline_model_parallel_size: number of GPUs used for pipeline model - parallelism. - attention_data_parallel_size: number of GPUs used for attention data - parallelism. - attention_context_model_parallel_size: number of GPUs used for attention context - parallelism. - moe_data_model_parallel_size: number of GPUs used for moe data - parallelism. - decode_context_parallel_size: number of GPUs used for decode context - parallelism, which splits the KV cache across GPUs within each - tensor-parallel group during decoding. Must be a divisor of - tensor_model_parallel_size and is currently only supported on the - AMD HIP platform. - shared_experts_tensor_parallel_size: optional shared-expert TP width. - Must divide attention TP; subgroups never cross attention replicas. + Every width comes from the runtime context rather than from an argument: + the configuration already says how wide each dimension is, and a caller + that translates it again is a second place for the two to disagree. A + process that needs a narrower layout than the one it published -- the + media encoder is the case in the tree -- states that layout on the context + first, so what it builds and what it answers stay the same thing. + + The remaining arguments are not topology. `backend` is decided by the + device, `duplicate_tp_group` and `enable_symm_mem` by other namespaces, and + `recovered_rank` / `rank_offset` / `max_world_size` describe this + particular join rather than the layout being joined. + + The widths this reads: + tp_size: GPUs used for tensor model parallelism. + moe_ep_size: GPUs used for expert model parallelism. + pp_size: GPUs used for pipeline model parallelism. + attn_dp_size: GPUs used for attention data parallelism. + attn_cp_size: GPUs used for attention context parallelism. + moe_dp_size: GPUs used for MoE data parallelism. + attn_dcp_size: GPUs used for decode context parallelism, which splits + the KV cache across GPUs within each tensor-parallel group during + decoding. Must be a divisor of `tp_size` and is currently only + supported on the AMD HIP platform. + shared_experts_tp_size: optional shared-expert TP width. Must divide + attention TP; subgroups never cross attention replicas. Let's say we have a total of 8 GPUs denoted by g0 ... g7 and we use 2 GPUs to parallelize the model tensor, and 4 GPUs to parallelize @@ -2589,6 +2587,16 @@ def initialize_model_parallel( assert torch.distributed.is_initialized() backend = backend or torch.distributed.get_backend(get_world_group().device_group) + parallel = get_parallel() + tensor_model_parallel_size = parallel.tp_size + expert_model_parallel_size = parallel.moe_ep_size + pipeline_model_parallel_size = parallel.pp_size + attention_data_parallel_size = parallel.attn_dp_size + attention_context_model_parallel_size = parallel.attn_cp_size + moe_data_model_parallel_size = parallel.moe_dp_size + decode_context_parallel_size = parallel.attn_dcp_size + shared_experts_tensor_parallel_size = parallel.shared_experts_tp_size + # Joiners construct their local TP/PP layout in global rank space. world_size: int = ( tensor_model_parallel_size * pipeline_model_parallel_size @@ -2933,6 +2941,12 @@ def initialize_model_parallel( max_world_size=max_world_size, ) + # The groups just built and the configuration they were built from are two + # accounts of one layout. Check them against each other here, where the + # disagreement is still attributable, rather than letting a collective run + # on the wrong peers. + _validate_parallel(get_parallel(), "group build") + def create_custom_parallel_group( group_ranks: List[int], backend: str = "gloo" @@ -3094,10 +3108,18 @@ def patch_tensor_parallel_group(tp_group: GroupCoordinator, *, owns_attention: b narrowed.update( attn_tp_size=tp_group.world_size, attn_tp_rank=tp_group.rank_in_group, + attn_tp_group=tp_group, attn_dp_size=1, attn_dp_rank=0, attn_cp_size=1, attn_cp_rank=0, + attn_cp_group=None, + moe_ep_size=1, + moe_ep_rank=0, + moe_ep_group=None, + moe_dp_size=1, + moe_tp_size=tp_group.world_size, + moe_tp_rank=tp_group.rank_in_group, ) try: with get_parallel().override(**narrowed): diff --git a/python/sglang/srt/layers/dp_attention.py b/python/sglang/srt/layers/dp_attention.py index 4919c34cd..c77d916bb 100644 --- a/python/sglang/srt/layers/dp_attention.py +++ b/python/sglang/srt/layers/dp_attention.py @@ -15,11 +15,9 @@ from sglang.srt.arg_groups.model_override_base import ( ) from sglang.srt.distributed import ( GroupCoordinator, - get_attn_tensor_model_parallel_world_size, ) from sglang.srt.distributed import get_moe_dp_group as _get_moe_dp_group from sglang.srt.distributed import ( - get_tensor_model_parallel_world_size, tensor_model_parallel_all_reduce, ) from sglang.srt.distributed.device_communicators.pynccl_allocator import ( @@ -49,6 +47,33 @@ if TYPE_CHECKING: from sglang.srt.model_executor.forward_batch_info import ForwardBatch +def dp_gather_width() -> int: + """How many replicas the DP sync gathers over. + + The attention-DP replicas, except after an elastic-EP scale-up, when the + gather spans the expanded WORLD -- whose width is the `dp_size` the + scale-up published. Read from the context either way: a scoped width has + to reach this, which is the whole reason the name has one home. + """ + parallel = get_parallel() + return parallel.dp_size if world_dp_gather_enabled() else parallel.attn_dp_size + + +def dp_gather_slot() -> int: + """This process's index in the list the DP sync just gathered. + + The gather spans the attention-DP replicas, except after an elastic-EP + scale-up, when it spans the expanded WORLD and the joining cohort is + numbered from its offset. Which list was gathered is what the flag below + says, so the index is read from there rather than kept as a second name on + the topology. + """ + parallel = get_parallel() + if world_dp_gather_enabled(): + return parallel.tp_rank + parallel.ep_join_rank_offset + return parallel.attn_dp_rank + + def world_dp_gather_enabled() -> bool: """Whether DP gathers should use expanded WORLD after joiner admission.""" dp = get_flags().dp @@ -60,9 +85,13 @@ def enable_joiner_all_gather(): def update_dp_attention_post_scale(new_dp_size: int, new_dp_rank: int): - get_parallel().override_permanently( - attn_dp_size=new_dp_size, attn_dp_rank=new_dp_rank - ) + """Point the DP gather at the expanded WORLD. + + The widths themselves are not written here: the caller scales `dp_size` on + the published bag, and the gather reads its width and this process's slot + from there. The arguments are the values the caller is about to publish, + kept so the log says which scale-up this was. + """ get_flags().dp.use_world_group_for_gather = True logger.debug( "[Elastic EP] dp_attention switched to WORLD: dp_size=%d dp_rank=%d", @@ -92,7 +121,7 @@ class DpPaddingMode(IntEnum): def get_dp_padding_mode( cls, is_extend_in_batch, global_num_tokens: List[int] ) -> DpPaddingMode: - dp_size = get_parallel().attn_dp_size + dp_size = dp_gather_width() # (trangdough) pplx-kernels a2a is a symmetric collective: every EP rank # must dispatch the same number of tokens or the device-side handshake @@ -374,14 +403,13 @@ def initialize_dp_attention( dp.enabled = enable_dp_attention tp_rank = get_parallel().tp_rank - tp_size = get_tensor_model_parallel_world_size() + tp_size = get_parallel().tp_size _, _, attn_dp_rank, attn_dp_size = compute_dp_attention_world_info( enable_dp_attention, tp_rank, tp_size, dp_size, attn_cp_size ) if get_exec().moe.elastic_ep_backend is not None and get_parallel().max_ep_size: - attn_dp_rank = tp_rank + get_parallel().ep_join_rank_offset # Reads the resolution, not a bag: this runs under # `initialize_dp_attention`, which the weight-cache daemon calls from # `_init_distributed` -- and other callers reach it from processes @@ -414,7 +442,10 @@ def is_allocation_symmetric() -> bool: def get_dp_local_info(forward_batch: ForwardBatch) -> Tuple[torch.Tensor, torch.Tensor]: # `get_dp_local_info` is only called in global DP gather and scatter. We use global DP rank here. - dp_rank = get_parallel().attn_dp_rank + # The slot in the list that was gathered. A scale-up widens that list + # to WORLD, and this process's index in it is not its index among the + # launch replicas. + dp_rank = dp_gather_slot() if forward_batch.dp_local_start_pos is None: cumtokens = torch.cumsum(forward_batch.global_num_tokens_gpu, dim=0) @@ -438,7 +469,10 @@ def get_dp_local_slice_cpu( # CPU (start, length) slice for DP-local data in a rank-padded buffer. # Returns Python ints (no D2H sync) and handles the cuda-graph-padded layout. global_num_tokens = forward_batch.global_num_tokens_cpu - dp_rank = get_parallel().attn_dp_rank + # The slot in the list that was gathered. A scale-up widens that list + # to WORLD, and this process's index in it is not its index among the + # launch replicas. + dp_rank = dp_gather_slot() local_num_tokens = global_num_tokens[dp_rank] if can_run_graph: local_start_pos = dp_rank * cuda_graph_batch @@ -514,7 +548,7 @@ def _dp_gather_via_all_reduce( NUM_GPUS_PER_NODE = 8 if ( not local_tokens.dtype.is_floating_point - and get_tensor_model_parallel_world_size() <= NUM_GPUS_PER_NODE + and get_parallel().tp_size <= NUM_GPUS_PER_NODE ): from sglang.srt.distributed.parallel_state import inplace_all_reduce @@ -534,7 +568,7 @@ def _dp_gather_via_all_gather( ): use_world = world_dp_gather_enabled() - if get_attn_tensor_model_parallel_world_size() == 1: + if get_parallel().attn_tp_size == 1: if use_world: torch.distributed.all_gather_into_tensor( global_tokens, @@ -548,9 +582,9 @@ def _dp_gather_via_all_gather( if not is_partial: if get_parallel().attn_tp_rank != 0: local_tokens.fill_(0) - scattered_local_tokens = local_tokens.tensor_split( - get_attn_tensor_model_parallel_world_size() - )[get_parallel().attn_tp_rank] + scattered_local_tokens = local_tokens.tensor_split(get_parallel().attn_tp_size)[ + get_parallel().attn_tp_rank + ] get_parallel().attn_tp_group.reduce_scatter_tensor( scattered_local_tokens, local_tokens ) @@ -721,8 +755,8 @@ def is_dp_gatherv_active() -> bool: return ( _USE_DP_GATHERV and not world_dp_gather_enabled() - and get_attn_tensor_model_parallel_world_size() == 1 - and get_tensor_model_parallel_world_size() == get_parallel().attn_dp_size + and get_parallel().attn_tp_size == 1 + and get_parallel().tp_size == get_parallel().attn_dp_size and not _DpGatheredBufferWrapper.is_dp_max_padding() ) @@ -881,12 +915,12 @@ def dp_reduce_scatter_tensor(output: torch.Tensor, input: torch.Tensor): if sizes is not None: get_parallel().tp_group.reduce_scatterv(input, output=output, sizes=sizes) return - if get_tensor_model_parallel_world_size() == get_parallel().attn_dp_size: + if get_parallel().tp_size == get_parallel().attn_dp_size: get_parallel().tp_group.reduce_scatter_tensor(output, input) else: - scattered_local_tokens = input.tensor_split( - get_tensor_model_parallel_world_size() - )[get_parallel().tp_rank] + scattered_local_tokens = input.tensor_split(get_parallel().tp_size)[ + get_parallel().tp_rank + ] get_parallel().tp_group.reduce_scatter_tensor(scattered_local_tokens, input) get_parallel().attn_tp_group.all_gather_into_tensor( output, scattered_local_tokens diff --git a/python/sglang/srt/layers/sampler.py b/python/sglang/srt/layers/sampler.py index de26d1eb6..341073e16 100644 --- a/python/sglang/srt/layers/sampler.py +++ b/python/sglang/srt/layers/sampler.py @@ -112,7 +112,12 @@ class Sampler(nn.Module): self.cp_sync_group = None if is_dp_attention_enabled(): self.tp_sync_group = get_parallel().attn_tp_group.device_group - self.cp_sync_group = get_parallel().attn_cp_group.device_group + # Only when there is more than one context shard to reconcile. The + # sync below already short-circuits on that, and a model running on + # one shard -- a speculative draft, under the scope that says so -- + # has no context-parallel communicator to name. + if get_parallel().attn_cp_size > 1: + self.cp_sync_group = get_parallel().attn_cp_group.device_group self.rl_on_policy_target = get_exec().deterministic.rl_on_policy_target # In RL on-policy mode, deterministic inference is automatically enabled. diff --git a/python/sglang/srt/managers/scheduler_components/dp_attn.py b/python/sglang/srt/managers/scheduler_components/dp_attn.py index a822d380e..1e924b179 100644 --- a/python/sglang/srt/managers/scheduler_components/dp_attn.py +++ b/python/sglang/srt/managers/scheduler_components/dp_attn.py @@ -10,7 +10,7 @@ from sglang.srt.configs.model_config import ModelConfig from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.environ import envs from sglang.srt.layers.cp.utils import get_cp_strategy -from sglang.srt.layers.dp_attention import world_dp_gather_enabled +from sglang.srt.layers.dp_attention import dp_gather_width, world_dp_gather_enabled from sglang.srt.layers.moe.utils import get_moe_a2a_backend from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.scheduler_components.recv_skipper import ( @@ -58,7 +58,7 @@ def _resolve_elastic_world_dp_size( from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager - live_dp_size = get_parallel().attn_dp_size + live_dp_size = dp_gather_width() effective_ep_size = ElasticEPStateManager.get_effective_ep_size() # The group's own membership, not the width it was built at: this is the # one number an out-of-process join moves, and it is the upper bound the diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 998a788cb..0bf416116 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -45,6 +45,7 @@ from sglang.srt.kv_canary.req_to_expected_token_ids_manager import ( ) from sglang.srt.layers.dp_attention import ( DpPaddingMode, + dp_gather_slot, set_dp_buffer_len, set_is_extend_in_batch, world_dp_gather_enabled, @@ -1165,9 +1166,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): if self.global_num_tokens_cpu is not None: # DP / MLP-sync path: per-DP padded width. if require_mlp_tp_gather(): - num_tokens_per_dp = self.global_num_tokens_cpu[ - get_parallel().attn_dp_rank - ] + num_tokens_per_dp = self.global_num_tokens_cpu[dp_gather_slot()] else: num_tokens_per_dp = self.global_num_tokens_cpu[0] else: @@ -1520,7 +1519,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): buffer_len = sum(global_num_tokens) if len(global_num_tokens) > 1: - num_tokens = global_num_tokens[get_parallel().attn_dp_rank] + num_tokens = global_num_tokens[dp_gather_slot()] else: num_tokens = global_num_tokens[0] diff --git a/python/sglang/srt/models/inkling_common/moe.py b/python/sglang/srt/models/inkling_common/moe.py index f5b6c52e5..b7644bb4c 100644 --- a/python/sglang/srt/models/inkling_common/moe.py +++ b/python/sglang/srt/models/inkling_common/moe.py @@ -673,9 +673,15 @@ class InklingSharedFusedMoE(FusedMoE): ) -> None: # FusedMoE.__init__ reads get_parallel() once and caches it on self, so # scoping the override to just this call is sufficient for the module's lifetime. + # The shared experts are replicated rather than sharded, so there is no + # expert-parallel communication here and no group to name: a width of + # one with the wider group still installed would describe a layout that + # does not exist. with get_parallel().override( moe_ep_size=1, moe_ep_rank=0, + moe_ep_group=None, + moe_dp_size=1, moe_tp_size=get_parallel().tp_size, moe_tp_rank=get_parallel().tp_rank, ): diff --git a/python/sglang/srt/models/kimi_k3.py b/python/sglang/srt/models/kimi_k3.py index a28f7b3fd..2bf00c929 100644 --- a/python/sglang/srt/models/kimi_k3.py +++ b/python/sglang/srt/models/kimi_k3.py @@ -22,7 +22,6 @@ from sglang.srt.configs.kimi_k3 import KimiK3Config from sglang.srt.configs.kimi_linear import KimiLinearConfig from sglang.srt.distributed import ( divide, - get_shared_experts_tp_group, tensor_model_parallel_all_reduce, ) from sglang.srt.distributed.device_communicators.pynccl_allocator import ( @@ -582,7 +581,7 @@ class KimiK3MoE(nn.Module): shared_experts_tp_kwargs = dict(tp_rank=0, tp_size=1) elif self._shared_experts_tp_comm: group = ( - get_shared_experts_tp_group() + parallel.shared_experts_tp_group if requested_shared_tp is not None else parallel.attn_tp_group ) diff --git a/python/sglang/srt/runtime_context.py b/python/sglang/srt/runtime_context.py index f8c424fd5..81736cabb 100644 --- a/python/sglang/srt/runtime_context.py +++ b/python/sglang/srt/runtime_context.py @@ -218,6 +218,7 @@ _LIVE_READS: dict = { "moe_tp_group": "get_moe_tp_group", "attn_tp_group": "get_attn_tp_group", "attn_cp_group": "get_attn_cp_group", + "shared_experts_tp_group": "get_shared_experts_tp_group", "dcp_group": "get_dcp_group", } @@ -448,14 +449,138 @@ class SpawnRanks(msgspec.Struct, frozen=True): the rank cannot say which replica this is. `None` means "no controller", which is an answer rather than an absence, and it is recorded as one. - Nothing else belongs here. A device index, for instance, is a placement - decision rather than a position -- the launcher may reindex it, and Ray - assigns it from its own allocator -- so it stays an argument to whoever - was handed it. + `gpu_id` is the device the parent picked for this process. It is not a + position in any group -- reindexing narrows the visible devices before the + spawn, and Ray allocates from its own pool -- but it is the same kind of + fact: something only the entry that spawned the process can state. `None` + for a process that runs on no device. """ world_rank: int dp_rank: Optional[int] = None + gpu_id: Optional[int] = None + + +_RANK_AND_WIDTH = ( + ("tp_rank", "tp_size"), + ("pp_rank", "pp_size"), + ("attn_tp_rank", "attn_tp_size"), + ("attn_dp_rank", "attn_dp_size"), + ("attn_cp_rank", "attn_cp_size"), + ("moe_ep_rank", "moe_ep_size"), +) + +# `moe_dp` is absent because `initialize_model_parallel` aliases the MoE-DP +# group to the attention-CP group when the latter is wider: there the group and +# the name are two facts, which is the same reason `moe_dp_rank` is left off the +# record at publish. +_WIDTH_AND_GROUP = ( + ("tp_size", "tp_group"), + ("pp_size", "pp_group"), + ("attn_tp_size", "attn_tp_group"), + ("attn_cp_size", "attn_cp_group"), + ("moe_ep_size", "moe_ep_group"), +) + +_UNREADABLE = object() + + +def _validate_parallel(parallel, source: str) -> None: + """Fail on a topology that cannot describe a real process layout. + + Every identity holds unconditionally: a width and a rank are both + plausible small integers whichever way they are wrong, so an inconsistent + set is not caught by anything downstream -- it surfaces as a hang or a + wrong answer in a collective, far from the write. Stating one leaf without + the quotients that follow from it leaves the namespace describing no real + layout, and the caller that did so is the one that has to say what it meant. + + Names that cannot be read are skipped rather than treated as zero: a + process that has published nothing can still stamp a rank, and a group that + has not been built answers nothing at all. + """ + + def read(name): + """A width or a rank, or `_UNREADABLE` for anything these identities + cannot be stated about -- an absent name, `None`, or a stand-in a test + put in a group's place. Booleans are integers in Python and are not + widths, so they are out too.""" + try: + value = getattr(parallel, name) + except Exception: + return _UNREADABLE + if isinstance(value, bool) or not isinstance(value, int): + return _UNREADABLE + return value + + problems = [] + + for rank_name, size_name in _RANK_AND_WIDTH: + rank, size = read(rank_name), read(size_name) + if _UNREADABLE in (rank, size): + continue + if not 0 <= rank < size: + problems.append( + f"0 <= {rank_name} < {size_name}\n {rank} is not a rank of {size}" + ) + + terms = ("tp_size", "attn_tp_size", "attn_dp_size", "attn_cp_size") + tp_size, a_tp, a_dp, a_cp = (read(n) for n in terms) + if _UNREADABLE not in (tp_size, a_tp, a_dp, a_cp): + if tp_size != a_tp * a_dp * a_cp: + problems.append( + "tp_size == attn_tp_size * attn_dp_size * attn_cp_size\n" + f" {tp_size} != {a_tp} * {a_dp} * {a_cp} (= {a_tp * a_dp * a_cp})" + ) + + moe_terms = ("tp_size", "moe_ep_size", "moe_dp_size", "moe_tp_size") + tp_size, m_ep, m_dp, m_tp = (read(n) for n in moe_terms) + if _UNREADABLE not in (tp_size, m_ep, m_dp, m_tp): + if tp_size != m_ep * m_dp * m_tp: + problems.append( + "tp_size == moe_ep_size * moe_dp_size * moe_tp_size\n" + f" {tp_size} != {m_ep} * {m_dp} * {m_tp} (= {m_ep * m_dp * m_tp})" + ) + + layout_terms = ( + "tp_rank", + "attn_dp_rank", + "attn_cp_rank", + "attn_tp_rank", + "attn_cp_size", + "attn_tp_size", + ) + tp_rank, r_dp, r_cp, r_tp, w_cp, w_tp = (read(n) for n in layout_terms) + if _UNREADABLE not in (tp_rank, r_dp, r_cp, r_tp, w_cp, w_tp): + laid_out = (r_dp * w_cp + r_cp) * w_tp + r_tp + if tp_rank != laid_out: + problems.append( + "tp_rank == (attn_dp_rank * attn_cp_size + attn_cp_rank)" + " * attn_tp_size + attn_tp_rank\n" + f" {tp_rank} != ({r_dp} * {w_cp} + {r_cp})" + f" * {w_tp} + {r_tp} (= {laid_out})" + ) + + for size_name, group_name in _WIDTH_AND_GROUP: + size = read(size_name) + if size is _UNREADABLE: + continue + try: + group = getattr(parallel, group_name) + except Exception: + continue + built = getattr(group, "world_size", _UNREADABLE) + if isinstance(built, int) and not isinstance(built, bool) and built != size: + problems.append( + f"{group_name}.world_size == {size_name}\n" + f" built {built}, configured {size}" + ) + + if problems: + raise ValueError( + f"parallel topology is inconsistent (set by {source}):\n" + + "\n".join(problems) + ) class ParallelContext: @@ -561,7 +686,13 @@ class ParallelContext: unknown = set(values) - _parallel_fields() if unknown: raise ValueError(f"unknown parallel field(s): {sorted(unknown)}") + saved = dict(self._stamp) self._stamp.update(values) + try: + _validate_parallel(self, "override_permanently") + except Exception: + self._stamp = saved + raise def clear_stamp(self) -> None: """Drop every stamped name, ranks included.""" @@ -576,6 +707,11 @@ class ParallelContext: raise ValueError(f"unknown parallel field(s): {sorted(unknown)}") saved = dict(self._overrides) self._overrides.update(kwargs) + try: + _validate_parallel(self, "override") + except Exception: + self._overrides = saved + raise try: yield self finally: @@ -1776,6 +1912,8 @@ def publish( ), ) _CONTEXT._publish_role = role + if ranks is not None and ranks.gpu_id is not None: + _CONTEXT.override("spawn", gpu_id=ranks.gpu_id) if ranks is not None: # The placement, worked out here rather than carried: the widths are on # the bag a moment ago, and `world_rank` fixes the rest. A read of any @@ -1804,8 +1942,13 @@ def publish( # "no controller" rather than an absence. placement["dp_rank"] = ranks.dp_rank placement["launch_world_rank"] = ranks.world_rank + placement.update(_attention_ranks(parallel, placement["tp_rank"])) + # One stamp, not two: the identities are checked on every write, and a + # half-placed process satisfies none of them. parallel.override_permanently(**placement) - _stamp_attention_ranks(parallel, placement["tp_rank"]) + # Publish established the whole layout, so every identity applies here, + # not just the ones the stamp happened to name. + _validate_parallel(parallel, "publish") if _ROLE_NS_MODE == "record": # The '-' marker distinguishes a zero-read role from a process where # recording never ran (signal teardown skips atexit). @@ -1820,7 +1963,7 @@ def publish( return _CONTEXT -def _stamp_attention_ranks(parallel, tp_rank: int) -> None: +def _attention_ranks(parallel, tp_rank: int) -> dict: """Place this process in the attention topology, from the configuration. The widths are already on the bag -- `publish` computed them a moment ago -- @@ -1829,8 +1972,9 @@ def _stamp_attention_ranks(parallel, tp_rank: int) -> None: that never initialises distributed, which is what `ParallelState` provided by being a plain frozen record. - It is a stamp rather than a bag leaf because it is a per-process fact, and - nothing about the configuration distinguishes one rank from another. + These are stamped rather than written as bag leaves because they are + per-process facts, and nothing about the configuration distinguishes one + rank from another. """ attn_tp_rank, attn_dp_rank = derive_attention_ranks( tp_rank=tp_rank, @@ -1838,7 +1982,7 @@ def _stamp_attention_ranks(parallel, tp_rank: int) -> None: attn_cp_size=parallel.attn_cp_size, enable_dp_attention=parallel.enable_dp_attention, ) - parallel.override_permanently(attn_tp_rank=attn_tp_rank, attn_dp_rank=attn_dp_rank) + return {"attn_tp_rank": attn_tp_rank, "attn_dp_rank": attn_dp_rank} def assert_published(server_args, *, role: str) -> RuntimeContext: diff --git a/python/sglang/srt/weight_cache/daemon.py b/python/sglang/srt/weight_cache/daemon.py index c68af59e2..0e7ae74fc 100644 --- a/python/sglang/srt/weight_cache/daemon.py +++ b/python/sglang/srt/weight_cache/daemon.py @@ -232,16 +232,7 @@ class WeightCacheDaemon: moe_a2a_backend=self.moe_a2a_backend, ) - initialize_model_parallel( - tensor_model_parallel_size=self.tp_size, - pipeline_model_parallel_size=self.pp_size, - expert_model_parallel_size=self.ep_size, - attention_data_parallel_size=( - self.dp_size if self.enable_dp_attention else 1 - ), - attention_context_model_parallel_size=self.attn_cp_size, - moe_data_model_parallel_size=self.moe_dp_size, - ) + initialize_model_parallel() # Initialize DP attention state (required by some models like Qwen3 MoE) from sglang.srt.layers.dp_attention import initialize_dp_attention diff --git a/python/sglang/test/layer_ut_utils.py b/python/sglang/test/layer_ut_utils.py index 01a9f24fc..bd324e73b 100644 --- a/python/sglang/test/layer_ut_utils.py +++ b/python/sglang/test/layer_ut_utils.py @@ -7,6 +7,8 @@ import os import torch +from sglang.test.test_utils import publish_build_topology + def init_single_process_dist(master_port: int = 29632, backend: str = "gloo"): """world=1 dist + model-parallel groups; srt layers require them even @@ -29,12 +31,8 @@ def init_single_process_dist(master_port: int = 29632, backend: str = "gloo"): if not model_parallel_is_initialized(): # kwargs only: a positional backend would land in the # attention_data_parallel_size slot and explode on int // str. - initialize_model_parallel( - tensor_model_parallel_size=1, - expert_model_parallel_size=1, - pipeline_model_parallel_size=1, - backend=backend, - ) + publish_build_topology(tp_size=1, ep_size=1, pp_size=1) + initialize_model_parallel(backend=backend) def make_tp1_column_parallel_linear( diff --git a/python/sglang/test/test_utils.py b/python/sglang/test/test_utils.py index 86165aaac..aae6098b0 100644 --- a/python/sglang/test/test_utils.py +++ b/python/sglang/test/test_utils.py @@ -2084,6 +2084,28 @@ def published_topology(role: str = "test", *, ranks=None, **server_args_fields): reset_context() +def publish_build_topology(*, world_rank: int = 0, **server_args_fields): + """State the widths `initialize_model_parallel` is about to build at. + + The build reads every width from the runtime context, so a test that wants + a particular topology publishes it here rather than passing it in -- the + same door production uses, which also keeps the derived widths honest. + + Unlike `published_topology` this is not a scope: the groups it is about to + build outlive any block, so the configuration describing them has to as + well. Callers that tear the groups down are already resetting the process. + """ + from sglang.srt.runtime_context import SpawnRanks, publish, reset_context + from sglang.srt.server_args import ServerArgs + + reset_context() + publish( + ServerArgs(model_path="dummy", **server_args_fields), + role="test", + ranks=SpawnRanks(world_rank=world_rank), + ) + + _GPU_IDLE_TIMEOUT_SECS = 30.0 _GPU_IDLE_POLL_INTERVAL_SECS = 2.0 _GPU_IDLE_USED_MEMORY_THRESHOLD = 2 << 30 # 2 GiB diff --git a/test/manual/ep/test_flashinfer_dispatcher.py b/test/manual/ep/test_flashinfer_dispatcher.py index 707dd5449..7aeb30cf5 100644 --- a/test/manual/ep/test_flashinfer_dispatcher.py +++ b/test/manual/ep/test_flashinfer_dispatcher.py @@ -14,7 +14,7 @@ from sglang.srt.layers.moe.token_dispatcher.flashinfer import FlashinferDispatch from sglang.srt.layers.moe.utils import initialize_moe_config from sglang.srt.runtime_context import get_context, publish from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler -from sglang.test.test_utils import CustomTestCase +from sglang.test.test_utils import CustomTestCase, publish_build_topology class TestFlashinferDispatcher(CustomTestCase): @@ -44,9 +44,8 @@ class TestFlashinferDispatcher(CustomTestCase): publish(server_args, role="scheduler") initialize_moe_config() - initialize_model_parallel( - tensor_model_parallel_size=world_size, expert_model_parallel_size=world_size - ) + publish_build_topology(tp_size=world_size, ep_size=world_size, world_rank=rank) + initialize_model_parallel() @classmethod def tearDownClass(cls): diff --git a/test/manual/test_custom_allreduce.py b/test/manual/test_custom_allreduce.py index 8261a36dd..e1d91d690 100644 --- a/test/manual/test_custom_allreduce.py +++ b/test/manual/test_custom_allreduce.py @@ -18,7 +18,7 @@ from sglang.srt.distributed.parallel_state import ( initialize_model_parallel, ) from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler -from sglang.test.test_utils import CustomTestCase +from sglang.test.test_utils import CustomTestCase, publish_build_topology def get_open_port() -> int: @@ -98,7 +98,8 @@ class TestCustomAllReduce(CustomTestCase): distributed_init_method=distributed_init_method, local_rank=rank, ) - initialize_model_parallel(tensor_model_parallel_size=world_size) + publish_build_topology(tp_size=world_size, world_rank=rank) + initialize_model_parallel() group = get_tensor_model_parallel_group().device_group # Set global server args to avoid "Global server args is not set yet!" error @@ -161,7 +162,8 @@ class TestCustomAllReduce(CustomTestCase): distributed_init_method=distributed_init_method, local_rank=rank, ) - initialize_model_parallel(tensor_model_parallel_size=world_size) + publish_build_topology(tp_size=world_size, world_rank=rank) + initialize_model_parallel() group = get_tensor_model_parallel_group().device_group # Set global server args to avoid "Global server args is not set yet!" error diff --git a/test/manual/test_quick_allreduce.py b/test/manual/test_quick_allreduce.py index f88049b03..3f4e3cd0b 100644 --- a/test/manual/test_quick_allreduce.py +++ b/test/manual/test_quick_allreduce.py @@ -23,7 +23,7 @@ from sglang.srt.distributed.parallel_state import ( graph_capture, initialize_model_parallel, ) -from sglang.test.test_utils import CustomTestCase +from sglang.test.test_utils import CustomTestCase, publish_build_topology torch.manual_seed(42) random.seed(44) # keep the deterministic seed @@ -117,7 +117,8 @@ class TestQuickAllReduce(CustomTestCase): distributed_init_method=distributed_init_method, local_rank=rank, ) - initialize_model_parallel(tensor_model_parallel_size=world_size) + publish_build_topology(tp_size=world_size, world_rank=rank) + initialize_model_parallel() group = get_tensor_model_parallel_group().device_group # A small all_reduce for warmup. @@ -186,7 +187,8 @@ class TestQuickAllReduce(CustomTestCase): distributed_init_method=distributed_init_method, local_rank=rank, ) - initialize_model_parallel(tensor_model_parallel_size=world_size) + publish_build_topology(tp_size=world_size, world_rank=rank) + initialize_model_parallel() group = get_tensor_model_parallel_group().device_group for sz in self.TEST_SIZES: diff --git a/test/manual/test_weight_loader_v2_equiv.py b/test/manual/test_weight_loader_v2_equiv.py index 1cf71c3e9..595b4339f 100644 --- a/test/manual/test_weight_loader_v2_equiv.py +++ b/test/manual/test_weight_loader_v2_equiv.py @@ -24,6 +24,7 @@ import unittest import torch from sglang.srt.environ import envs +from sglang.test.test_utils import publish_build_topology MODEL = "Qwen/Qwen2-0.5B" @@ -43,7 +44,8 @@ def _init_model_parallel() -> None: local_rank=0, distributed_init_method="tcp://127.0.0.1:29634", ) - initialize_model_parallel(tensor_model_parallel_size=1) + publish_build_topology(tp_size=1) + initialize_model_parallel() monkey_patch_vllm_parallel_state() except AssertionError: pass diff --git a/test/registered/cpu/test_rank_consensus_checker.py b/test/registered/cpu/test_rank_consensus_checker.py index 773c4b3ff..ebd183f5c 100644 --- a/test/registered/cpu/test_rank_consensus_checker.py +++ b/test/registered/cpu/test_rank_consensus_checker.py @@ -23,7 +23,11 @@ from sglang.srt.utils.rank_consensus_checker import ( shutdown, ) from sglang.test.ci.ci_register import register_cpu_ci -from sglang.test.test_utils import CustomTestCase, find_available_port +from sglang.test.test_utils import ( + CustomTestCase, + find_available_port, + publish_build_topology, +) register_cpu_ci(est_time=193, suite="stage-a-test-cpu-intel") @@ -80,11 +84,8 @@ def run_distributed_test( backend="gloo", ) - initialize_model_parallel( - tensor_model_parallel_size=tp_size, - pipeline_model_parallel_size=pp_size, - backend="gloo", - ) + publish_build_topology(tp_size=tp_size, pp_size=pp_size, world_rank=rank) + initialize_model_parallel(backend="gloo") fn() except Exception as e: diff --git a/test/registered/e2e/pp/test_pp_spec_embed_scan.py b/test/registered/e2e/pp/test_pp_spec_embed_scan.py index 5c8cd9c18..f42b5f85f 100644 --- a/test/registered/e2e/pp/test_pp_spec_embed_scan.py +++ b/test/registered/e2e/pp/test_pp_spec_embed_scan.py @@ -10,7 +10,11 @@ import torch from transformers import MistralConfig, PretrainedConfig from sglang.test.ci.ci_register import register_cuda_ci -from sglang.test.test_utils import DEFAULT_SMALL_MODEL_NAME_FOR_TEST, CustomTestCase +from sglang.test.test_utils import ( + DEFAULT_SMALL_MODEL_NAME_FOR_TEST, + CustomTestCase, + publish_build_topology, +) register_cuda_ci(est_time=60, stage="base-b", runner_config="1-gpu-small") @@ -77,7 +81,8 @@ class TestDraftEmbedScan(CustomTestCase): init_distributed_environment( world_size=1, rank=0, local_rank=0, distributed_init_method="env://" ) - initialize_model_parallel(tensor_model_parallel_size=1) + publish_build_topology(tp_size=1) + initialize_model_parallel() torch.set_default_dtype(torch.bfloat16) torch.cuda.set_device(0) diff --git a/test/registered/eplb/test_lplb_distributed.py b/test/registered/eplb/test_lplb_distributed.py index 59edf9300..5d9220d61 100644 --- a/test/registered/eplb/test_lplb_distributed.py +++ b/test/registered/eplb/test_lplb_distributed.py @@ -38,6 +38,7 @@ from sglang.srt.distributed.parallel_state import ( initialize_model_parallel, ) from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import publish_build_topology register_cuda_ci(est_time=18, stage="base-b", runner_config="2-gpu-large") @@ -238,10 +239,10 @@ def _worker_main(local_rank: int, world_size: int): init_distributed_environment( world_size=world_size, rank=local_rank, local_rank=local_rank ) - initialize_model_parallel( - tensor_model_parallel_size=world_size, - expert_model_parallel_size=world_size, + publish_build_topology( + tp_size=world_size, ep_size=world_size, world_rank=local_rank ) + initialize_model_parallel() from sglang.srt.eplb.lplb_solver import clear_global_lplb_solvers diff --git a/test/registered/kernels/ops/communication/test_deterministic_reduce_scatter.py b/test/registered/kernels/ops/communication/test_deterministic_reduce_scatter.py index 703718857..9eda92e6c 100644 --- a/test/registered/kernels/ops/communication/test_deterministic_reduce_scatter.py +++ b/test/registered/kernels/ops/communication/test_deterministic_reduce_scatter.py @@ -14,6 +14,7 @@ from sglang.srt.distributed import parallel_state as ps from sglang.srt.environ import envs from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.kernels.utils import multigpu_pytest_main +from sglang.test.test_utils import publish_build_topology register_cuda_ci(est_time=45, stage="base-b-kernel-unit", runner_config="4-gpu-b200") @@ -37,7 +38,8 @@ def group(): local_rank=local_rank, distributed_init_method="env://", ) - ps.initialize_model_parallel(tensor_model_parallel_size=world_size) + publish_build_topology(tp_size=world_size, world_rank=rank) + ps.initialize_model_parallel() yield ps.get_tp_group() ps.destroy_model_parallel() ps.destroy_distributed_environment() diff --git a/test/registered/kernels/ops/moe/test_fused_swiglu_epilogue.py b/test/registered/kernels/ops/moe/test_fused_swiglu_epilogue.py index 51e7b2340..9e483f1ff 100644 --- a/test/registered/kernels/ops/moe/test_fused_swiglu_epilogue.py +++ b/test/registered/kernels/ops/moe/test_fused_swiglu_epilogue.py @@ -19,6 +19,7 @@ import pytest import torch from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import publish_build_topology register_cuda_ci(est_time=12, stage="base-b-kernel-unit", runner_config="1-gpu-large") @@ -52,12 +53,8 @@ def _runtime_scaffolding(): if not torch.distributed.is_initialized(): init_distributed_environment(world_size=1, rank=0, local_rank=0, backend="gloo") if not model_parallel_is_initialized(): - initialize_model_parallel( - tensor_model_parallel_size=1, - expert_model_parallel_size=1, - pipeline_model_parallel_size=1, - backend="gloo", - ) + publish_build_topology(tp_size=1, ep_size=1, pp_size=1) + initialize_model_parallel(backend="gloo") def _interleave_w13_rows(w13: torch.Tensor) -> torch.Tensor: diff --git a/test/registered/layers/mamba/test_mamba2_mixer.py b/test/registered/layers/mamba/test_mamba2_mixer.py index 063b84f37..7ac76d0c4 100644 --- a/test/registered/layers/mamba/test_mamba2_mixer.py +++ b/test/registered/layers/mamba/test_mamba2_mixer.py @@ -17,6 +17,7 @@ from sglang.srt.distributed.parallel_state import ( from sglang.srt.runtime_context import get_parallel from sglang.srt.utils import get_device, get_device_count from sglang.test.ci.ci_register import register_cuda_ci, register_xpu_ci +from sglang.test.test_utils import publish_build_topology register_cuda_ci(est_time=30, stage="base-b", runner_config="2-gpu-large") register_xpu_ci(est_time=60, suite="nightly-xpu-2-gpu", nightly=True) @@ -105,7 +106,8 @@ def mixer2_gated_norm_tensor_parallel( local_rank=local_rank, backend=get_default_distributed_backend(device.type), ) - initialize_model_parallel(tensor_model_parallel_size=world_size) + publish_build_topology(tp_size=world_size, world_rank=local_rank) + initialize_model_parallel() # create random weights an inputs weight = torch.rand((hidden_size,), dtype=dtype, device=device) diff --git a/test/registered/ops/test_aiter_allreduce_fusion_amd.py b/test/registered/ops/test_aiter_allreduce_fusion_amd.py index 926e4f338..8f3edf632 100755 --- a/test/registered/ops/test_aiter_allreduce_fusion_amd.py +++ b/test/registered/ops/test_aiter_allreduce_fusion_amd.py @@ -14,7 +14,7 @@ import torch from sglang.srt.layers import communicator as comm from sglang.srt.layers.communicator import LayerCommunicator, ScatterMode from sglang.test.ci.ci_register import register_amd_ci -from sglang.test.test_utils import CustomTestCase +from sglang.test.test_utils import CustomTestCase, publish_build_topology register_amd_ci(est_time=240, suite="stage-c-test-large-8-gpu-amd") @@ -64,7 +64,8 @@ def _run_residual_accuracy_check(): distributed_init_method="env://", backend="nccl", ) - initialize_model_parallel(tensor_model_parallel_size=world_size) + publish_build_topology(tp_size=world_size, world_rank=rank) + initialize_model_parallel() dtype = torch.bfloat16 eps = 1e-6 diff --git a/test/registered/unit/distributed/test_parallel_state.py b/test/registered/unit/distributed/test_parallel_state.py index f870b19e0..edb86d654 100644 --- a/test/registered/unit/distributed/test_parallel_state.py +++ b/test/registered/unit/distributed/test_parallel_state.py @@ -44,6 +44,7 @@ import pytest import torch from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import publish_build_topology register_cpu_ci(est_time=11, suite="base-a-test-cpu") @@ -232,11 +233,8 @@ def test_parallel_group_construction_tp8_attn_cp2(): mock_world_group.return_value = mock_world # Call the actual function - parallel_state.initialize_model_parallel( - tensor_model_parallel_size=8, - pipeline_model_parallel_size=1, - attention_context_model_parallel_size=2, - ) + publish_build_topology(tp_size=8, pp_size=1, attn_cp_size=2) + parallel_state.initialize_model_parallel() # Verify TP groups tp_groups = created_groups.get("tp", []) @@ -330,12 +328,8 @@ def test_parallel_group_construction_tp8_moe_ep4_cp2(): mock_world_group.return_value = mock_world # Call the actual function - parallel_state.initialize_model_parallel( - tensor_model_parallel_size=8, - expert_model_parallel_size=4, - pipeline_model_parallel_size=1, - moe_data_model_parallel_size=2, - ) + publish_build_topology(tp_size=8, ep_size=4, pp_size=1, moe_dp_size=2) + parallel_state.initialize_model_parallel() # Verify TP groups tp_groups = created_groups.get("tp", []) diff --git a/test/registered/unit/layers/attention/test_vision_strided_qkv.py b/test/registered/unit/layers/attention/test_vision_strided_qkv.py index 380e542df..666d8d6a6 100644 --- a/test/registered/unit/layers/attention/test_vision_strided_qkv.py +++ b/test/registered/unit/layers/attention/test_vision_strided_qkv.py @@ -16,6 +16,7 @@ from sglang.srt.distributed.parallel_state import ( from sglang.srt.layers.attention import vision from sglang.srt.runtime_context import get_context, get_parallel from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import publish_build_topology register_cpu_ci(est_time=12, suite="base-a-test-cpu") @@ -52,7 +53,8 @@ def gloo_world(): distributed_init_method=f"tcp://127.0.0.1:{port}", backend="gloo", ) - initialize_model_parallel(tensor_model_parallel_size=1, backend="gloo") + publish_build_topology(tp_size=1) + initialize_model_parallel(backend="gloo") yield destroy_model_parallel() destroy_distributed_environment() diff --git a/test/registered/unit/mem_cache/test_dsa_layer_split_broadcast.py b/test/registered/unit/mem_cache/test_dsa_layer_split_broadcast.py index 7d2d49e93..209b7cd49 100644 --- a/test/registered/unit/mem_cache/test_dsa_layer_split_broadcast.py +++ b/test/registered/unit/mem_cache/test_dsa_layer_split_broadcast.py @@ -20,7 +20,7 @@ import torch import torch.multiprocessing as mp from sglang.test.ci.ci_register import register_cuda_ci -from sglang.test.test_utils import CustomTestCase +from sglang.test.test_utils import CustomTestCase, publish_build_topology register_cuda_ci(est_time=28, stage="base-c", runner_config="4-gpu-b200") @@ -54,10 +54,8 @@ def _run(rank: int, world: int, port: int): distributed_init_method=f"tcp://127.0.0.1:{port}", backend="nccl", ) - initialize_model_parallel( - tensor_model_parallel_size=world, - attention_context_model_parallel_size=world, - ) + publish_build_topology(tp_size=world, attn_cp_size=world, world_rank=rank) + initialize_model_parallel() from sglang.srt.mem_cache.dsa_cache_layer_split import ( LayerSplitDSATokenToKVPool, diff --git a/test/registered/unit/mem_cache/test_page_interleave_shard.py b/test/registered/unit/mem_cache/test_page_interleave_shard.py index 6961f75fc..8ddcd37d4 100644 --- a/test/registered/unit/mem_cache/test_page_interleave_shard.py +++ b/test/registered/unit/mem_cache/test_page_interleave_shard.py @@ -84,7 +84,7 @@ from sglang.srt.runtime_context import get_parallel, publish from sglang.srt.server_args import ServerArgs from sglang.srt.utils import ceil_div from sglang.test.ci.ci_register import register_cpu_ci -from sglang.test.test_utils import CustomTestCase +from sglang.test.test_utils import CustomTestCase, publish_build_topology register_cpu_ci(est_time=30, suite="base-a-test-cpu") @@ -1215,10 +1215,8 @@ def _dist_init(rank, world, port, attn_cp_size): ServerArgs(model_path="dummy", tp_size=world, attn_cp_size=attn_cp_size), role="scheduler", ) - initialize_model_parallel( - tensor_model_parallel_size=world, - attention_context_model_parallel_size=attn_cp_size, - ) + publish_build_topology(tp_size=world, attn_cp_size=attn_cp_size, world_rank=rank) + initialize_model_parallel() def _gather_make_spec(shard_rank, max_prefix_groups=16, chunk_groups=4): diff --git a/test/registered/unit/model_loader/test_presharded_loader.py b/test/registered/unit/model_loader/test_presharded_loader.py index bc74cd311..527259ee7 100644 --- a/test/registered/unit/model_loader/test_presharded_loader.py +++ b/test/registered/unit/model_loader/test_presharded_loader.py @@ -844,7 +844,9 @@ class TestShardConfig(unittest.TestCase): override.install() self.addCleanup(override.restore) with ( - get_parallel().override(tp_size=8, pp_size=1, moe_dp_size=2, moe_ep_size=4), + get_parallel().override( + tp_size=8, pp_size=1, moe_dp_size=2, moe_ep_size=4, moe_tp_size=1 + ), mock.patch( "sglang.srt.layers.dp_attention.get_moe_cp_size", return_value=2, diff --git a/test/registered/unit/models/test_glm5_next_bfg_fusion.py b/test/registered/unit/models/test_glm5_next_bfg_fusion.py index c7607ff6b..272c52b1e 100644 --- a/test/registered/unit/models/test_glm5_next_bfg_fusion.py +++ b/test/registered/unit/models/test_glm5_next_bfg_fusion.py @@ -83,8 +83,19 @@ class TestGlm5NextBfgFusion(unittest.TestCase): for attn_tp, rank in ((1, 0), (2, 0), (2, 1)): with ( self.subTest(route=expected_route, attn_tp=attn_tp, rank=rank), + # A width is a whole topology: the attention triple has + # to factor `tp_size`, and this process has to sit where + # the triple puts it. get_parallel().override( - tp_size=4, tp_rank=3, attn_tp_size=attn_tp, attn_tp_rank=rank + tp_size=4, + tp_rank=rank, + attn_tp_size=attn_tp, + attn_tp_rank=rank, + attn_dp_size=4 // attn_tp, + attn_dp_rank=0, + attn_cp_size=1, + attn_cp_rank=0, + moe_tp_size=4, ), ): quant = MockFp8Config(ignored) diff --git a/test/registered/unit/models/test_shared_experts_fusion_gates.py b/test/registered/unit/models/test_shared_experts_fusion_gates.py index 7afae42ca..96f6dbe4b 100644 --- a/test/registered/unit/models/test_shared_experts_fusion_gates.py +++ b/test/registered/unit/models/test_shared_experts_fusion_gates.py @@ -61,7 +61,15 @@ class _FusionGateCase(CustomTestCase): def _reason(self, model_class, hf_config, quant_config=None, moe_ep_size=1): # The gates consult the live EP size; without a group installed the # canonical getter asserts, so every case states a topology. - with get_parallel().override(moe_ep_size=moe_ep_size): + with get_parallel().override( + tp_size=moe_ep_size, + attn_tp_size=moe_ep_size, + attn_dp_size=1, + attn_cp_size=1, + moe_ep_size=moe_ep_size, + moe_dp_size=1, + moe_tp_size=1, + ): return model_class.shared_experts_fusion_disable_reason( hf_config, quant_config ) diff --git a/test/registered/unit/test_runtime_context.py b/test/registered/unit/test_runtime_context.py index 27e9876f2..6c50ff3c1 100644 --- a/test/registered/unit/test_runtime_context.py +++ b/test/registered/unit/test_runtime_context.py @@ -41,6 +41,7 @@ from sglang.srt.runtime_context import ( RuntimeContext, SpawnRanks, _FlagGroupBase, + _validate_parallel, assert_published, derive_parallel_widths, get_context, @@ -59,6 +60,56 @@ from sglang.srt.server_args import ServerArgs from sglang.test.test_utils import CustomTestCase _SRT = _pathlib.Path(next(iter(_sglang.__path__))).resolve() / "srt" +_PACKAGE = _pathlib.Path(next(iter(_sglang.__path__))).resolve() + + +def _sources(): + """Every Python file this checkout ships, package and siblings alike. + + The package alone is the wrong subject set for anything about entries or + public names: `benchmark/`, `examples/` and the top-level `test/` call the + same doors and are not covered by any suite that would notice them break. + An installed package has no siblings, and then this is the package alone.""" + roots = [_PACKAGE] + checkout = _PACKAGE.parents[1] + roots += [ + checkout / name + for name in ("benchmark", "examples", "scripts", "test") + if (checkout / name).is_dir() + ] + for root in roots: + for path in root.rglob("*.py"): + yield path + + +def _scope_entries_that_say_nothing(paths): + """Draft-scope entries that do not state `owns_attention`, as `path:line`. + + The scope either narrows the draft's attention and expert identity or + leaves the target's in place, and only the worker knows which -- so the + keyword has no default. Omitting it is a `TypeError`, but only on the path + that runs, and those paths want a GPU and a draft model. + """ + import ast + + missing = [] + for path in paths: + try: + tree = ast.parse(path.read_text(encoding="utf-8")) + except (SyntaxError, UnicodeDecodeError): + continue + for node in ast.walk(tree): + if not isinstance(node, ast.Call): + continue + func = node.func + name = getattr(func, "attr", None) or getattr(func, "id", None) + if name not in ("draft_tp_context", "patch_tensor_parallel_group"): + continue + if not any(kw.arg == "owns_attention" for kw in node.keywords): + missing.append(f"{path}:{node.lineno}") + return missing + + _PS = "sglang.srt.distributed.parallel_state" _DP = "sglang.srt.layers.dp_attention" @@ -347,9 +398,8 @@ class TestStampedRanks(_IsolatedOverrides): """`attn_dp_rank` comes from the stamp, and says so when there is none. It is the one rank no group answers with: `initialize_dp_attention` - computes it from this process's `tp_rank`, and an elastic scale-up - replaces it with a rank in the expanded WORLD. Falling back to anything - would be inventing a placement for this process. + computes it from this process's `tp_rank`. Falling back to anything would + be inventing a placement for this process. """ def setUp(self): @@ -407,21 +457,71 @@ class TestStampedRanks(_IsolatedOverrides): ) self.assertIs(mode, DpPaddingMode.SUM_LEN) - def test_a_scale_up_stamps_the_width_and_the_rank_together(self): - """The two describe one topology; a reader that saw only one moved - would place this process in a group it is not in.""" + def test_the_gather_slot_follows_the_list_that_was_gathered(self): + """The DP sync gathers over the attention-DP replicas, or over the + expanded WORLD once a scale-up has moved the gather there. The index + into that list is a property of the gather, so it is read beside the + flag that says which one happened rather than kept on the topology.""" + from sglang.srt.layers.dp_attention import dp_gather_slot + + self.addCleanup(reset_context) + dp_flags = get_flags().dp + saved = ( + dp_flags.use_world_group_for_gather, + dp_flags.joiner_skip_all_gather, + ) + + def restore(): + ( + dp_flags.use_world_group_for_gather, + dp_flags.joiner_skip_all_gather, + ) = saved + + self.addCleanup(restore) + publish( + ServerArgs( + model_path="dummy", tp_size=8, dp_size=8, enable_dp_attention=True + ), + role="test", + ranks=SpawnRanks(world_rank=3), + ) + parallel = get_parallel() + dp_flags.use_world_group_for_gather = False + self.assertEqual(dp_gather_slot(), parallel.attn_dp_rank) + + # After a scale-up the gather spans the expanded WORLD, and the joining + # cohort is numbered from its offset. + dp_flags.use_world_group_for_gather = True + dp_flags.joiner_skip_all_gather = False + parallel.override_permanently(ep_join_rank_offset=8) + self.assertEqual(dp_gather_slot(), 8 + parallel.tp_rank) + # and the topology it was read off is untouched + self.assertEqual(parallel.attn_dp_size, 8) + self.assertEqual(parallel.tp_size, 8) + + def test_a_scale_up_writes_no_width(self): + """The identities stay unconditional because nothing overrides them: + the scale-up only points the gather at the expanded WORLD.""" from sglang.srt.layers.dp_attention import update_dp_attention_post_scale - # It also flips a process-wide gather flag; put it back, or every - # later test in this process runs as if a scale-up had happened. dp_flags = get_flags().dp saved_gather = dp_flags.use_world_group_for_gather self.addCleanup(setattr, dp_flags, "use_world_group_for_gather", saved_gather) - + self.addCleanup(reset_context) + publish( + ServerArgs( + model_path="dummy", tp_size=8, dp_size=8, enable_dp_attention=True + ), + role="test", + ranks=SpawnRanks(world_rank=3), + ) parallel = get_parallel() + before = (parallel.attn_dp_size, parallel.attn_dp_rank, parallel.tp_size) update_dp_attention_post_scale(new_dp_size=16, new_dp_rank=11) - self.assertEqual(parallel.attn_dp_size, 16) - self.assertEqual(parallel.attn_dp_rank, 11) + self.assertTrue(dp_flags.use_world_group_for_gather) + self.assertEqual( + (parallel.attn_dp_size, parallel.attn_dp_rank, parallel.tp_size), before + ) class TestEveryDeclaredParallelNameIsStatable(_IsolatedOverrides): @@ -1898,14 +1998,26 @@ class TestDerivedWidths(_IsolatedOverrides): def test_a_topology_is_stated_by_naming_the_width(self): """Overriding a leaf does not move the quotient -- the quotient is not - recomputed on read. Naming it is how a test states one.""" + recomputed on read. Naming it is how a caller states one, and naming + only some of them is refused: the caller owns the arithmetic, the + context only checks it.""" reset_context() self.addCleanup(reset_context) publish(ServerArgs(model_path="dummy", tp_size=8), role="test") self.assertEqual(get_parallel().attn_tp_size, 8) - with get_parallel().override(tp_size=2): - self.assertEqual(get_parallel().attn_tp_size, 8) - with get_parallel().override(attn_tp_size=4): + + with self.assertRaises(ValueError) as caught: + with get_parallel().override(tp_size=2): + pass + self.assertIn( + "tp_size == attn_tp_size * attn_dp_size * attn_cp_size", + str(caught.exception), + ) + + with get_parallel().override(tp_size=2, attn_tp_size=2, moe_tp_size=2): + self.assertEqual(get_parallel().tp_size, 2) + self.assertEqual(get_parallel().attn_tp_size, 2) + with get_parallel().override(attn_tp_size=4, tp_size=4, moe_tp_size=4): self.assertEqual(get_parallel().attn_tp_size, 4) def test_an_unstated_topology_still_fails(self): @@ -2127,17 +2239,11 @@ class TestDerivedWidths(_IsolatedOverrides): ) self.assertEqual(published, recomputed) - def test_initialize_model_parallel_no_longer_touches_the_bag(self): - """`initialize_model_parallel` used to recompute and - permanently override the six derived widths on `get_parallel()` - after building its groups; that call is gone. Publish a placeholder - config (tp_size defaults to 1), then build real groups at a - different width -- the published leaf must now stay exactly what it - was, because nothing corrects it. This is the behavior a caller - relies on being told about, loudly, the first time it publishes and - builds inconsistently -- see - `test_recomputing_from_published_leaves_matches_the_publish_bag` - for why every real caller must not do that. + def test_initialize_model_parallel_builds_at_the_published_widths(self): + """The build takes every width from the context rather than from an + argument, so "published one width, built another" is no longer a state + a caller can reach -- there is nothing left to translate, and nothing + to correct afterwards either. """ from unittest.mock import Mock @@ -2145,11 +2251,11 @@ class TestDerivedWidths(_IsolatedOverrides): reset_context() self.addCleanup(reset_context) - publish(ServerArgs(model_path="dummy"), role="test") - self.assertEqual(get_parallel().attn_tp_size, 1) - self.assertEqual(get_parallel().moe_ep_size, 1) - world_size = 8 + publish(ServerArgs(model_path="dummy", tp_size=world_size), role="test") + self.assertEqual(get_parallel().attn_tp_size, world_size) + + built_at = [] with ( patch.object(parallel_state, "_WORLD", None), patch.object(parallel_state, "_TP", None), @@ -2168,25 +2274,20 @@ class TestDerivedWidths(_IsolatedOverrides): patch.object( parallel_state, "init_model_parallel_group", - return_value=Mock(device_group=Mock()), + side_effect=lambda group_ranks, *a, **k: ( + built_at.append(group_ranks), + Mock(device_group=Mock()), + )[1], ), patch.object(parallel_state, "get_world_group") as mock_world_group, ): mock_world_group.return_value = Mock(device_group=Mock(), local_rank=0) - parallel_state.initialize_model_parallel( - tensor_model_parallel_size=world_size, - expert_model_parallel_size=world_size, - ) + parallel_state.initialize_model_parallel() self.addCleanup(parallel_state.destroy_model_parallel) - self.assertEqual( - get_parallel().attn_tp_size, - 1, - "initialize_model_parallel must not touch the published leaf -- " - "a caller that needs it corrected must publish a config that " - "already matches the width it is about to build", - ) - self.assertEqual(get_parallel().moe_ep_size, 1) + # The first group built is TP, one group spanning the published width. + self.assertEqual(built_at[0], [list(range(world_size))]) + self.assertEqual(get_parallel().attn_tp_size, world_size) class TestTheDerivedHalfIsDeclared(CustomTestCase): @@ -2294,6 +2395,225 @@ class TestAnEntryThatBuildsARunnerHandsOverItsPlacement(CustomTestCase): ) +class TestTheAccessorsHaveNoCallersOutsideTheirPackage(CustomTestCase): + """`parallel_state`'s getters are the definition, not a second spelling. + + Business code asks `get_parallel()`; a call that goes straight to the getter + is a read the context cannot redirect, which is what a scope needs it to be + able to do. The package that defines them is exempt -- a read there would + go through the context back into itself -- and so is `multimodal_gen`, which + has its own parallel state. + """ + + #: Not topology. `get_self_pp_group` builds the single-rank group a draft + #: pipeline scope installs, so there is nothing for the context to answer + #: with until the scope has installed it. + ALLOWED = { + "get_self_pp_group", + "get_default_distributed_backend", + "get_mooncake_transfer_engine", + } + + def _accessors(self): + """Derived from the source, not listed here: a guard whose subject set + is written by hand stops watching whatever gets added next.""" + from sglang.srt.distributed import parallel_state as parallel_state_module + + source = _pathlib.Path(parallel_state_module.__file__).read_text().splitlines() + return { + line[len("def ") : line.index("(")] + for line in source + if line.startswith("def get_") or line.startswith("def is_") + } + + def _callers(self, name): + import re + + from sglang.srt.distributed import parallel_state as parallel_state_module + + root = _pathlib.Path(parallel_state_module.__file__).parents[2] + pattern = re.compile(rf"(?