diff --git a/.claude/skills/sglang-runtime-context/SKILL.md b/.claude/skills/sglang-runtime-context/SKILL.md index 0964e2dae09d..16de40f65fd1 100644 --- a/.claude/skills/sglang-runtime-context/SKILL.md +++ b/.claude/skills/sglang-runtime-context/SKILL.md @@ -15,7 +15,7 @@ One container owns process-static runtime state: `sglang.srt.runtime_context.Run | runtime flags | `get_flags()` | state that is *not* a pure function of config: `capture` (cuda-graph lifecycle), `moe` (ACTIVE backends, swappable), `dp` (DP-attention runtime flags) | materialized at subsystem init; groups offer `override()` for tests | | resources | `get_resources()`, `get_stream(name)`, `get_buffer(name, factory)` | process-level handles: graph pools, EPLB state, EP dispatcher state, named side streams, workspace buffers | lazy; cleared by `reset_context()` | | per-forward | `get_forward()` | forward-scoped flags (multi-stream switch, MoE output buffer, attn-TP inputs, extend-in-batch) | contextvar-backed; `scoped(**kw)` restores on exit; new threads see defaults | -| parallel | `get_parallel()` | **dual, spelled**: bare names are the live topology (tp/pp/moe/attn sizes, ranks, groups — `@property`, read-through); `get_parallel().config.` is the parallel config bag | live: after dist init; `config`: after publish | +| parallel | `get_parallel()` | one spelling per name: ranks and group handles are the live topology (`@property`, read-through); every other name, sizes included, is a leaf of the parallel config bag | ranks/groups: after dist init; leaves: after publish | `reset_context()` (unit-test teardown) drops the published config and installs fresh flags/resources/forward tiers. @@ -221,26 +221,33 @@ row for a method the Ray actor does not have, and an effective-field set without `load_format` -- and both were invisible because the assertion had slack (`>= len(...) - 1`) or compared key names instead of value sources. -### `get_parallel()`: live topology bare, configuration under `config` - -**Bare is the live group, `config` is what was configured.** `get_parallel().tp_size` -and its size / rank / group siblings are `@property` read-through over the canonical -getters; `get_parallel().config.` reads the published `parallel` bag -(`nccl_port`, `enable_dp_attention`, `dp_size`, `ep_size`, `dwdp_size`, ... and the -five sizes that also have a live property). A bare read of a config-only leaf raises -an `AttributeError` naming the `.config` spelling — the tier is never guessed from -whether a property happens to exist. - -The two tiers are **not** two spellings of one number. Live diverges from configured -wherever elastic EP scales the world away from the launch shape, and wherever -`initialize_model_parallel` aliases `_MOE_DP` to `_ATTN_CP` (`attn_cp_size > -moe_dp_size`), which makes a live comparison of that pair degenerate. The five -live-shadowed sizes (`tp/pp/dcp/attn_cp/moe_dp_size`) are where the choice matters, -and every business read of `get_parallel().config.` is registered with -its reason in `_CONFIGURED_SIZE_CALL_SITES` (`test_global_config_read_ratchet.py`). -DCP has a third shape: the live `get_parallel().attn_dcp_size` / `.dcp_enabled` -answer the *effective* topology (`1` / `False` with no group installed), never the -requested size — `.config.dcp_size` is the requested one. +### `get_parallel()`: one spelling per name + +**There is no `.config` hop.** Ranks and group handles are `@property` +read-through over the canonical getters, so they answer with the live process +groups. Everything else — `tp_size`, `pp_size`, `attn_cp_size`, `dcp_size`, +`moe_dp_size` included, alongside config-only leaves like `nccl_port`, +`enable_dp_attention`, `dp_size`, `ep_size`, `dwdp_size` — is answered from the +published `parallel` bag. Reading a leaf before publish raises a `ValueError` +naming the namespace; an unknown name is an `AttributeError`. + +A size reads from the configuration because the groups are built at exactly the +configured widths — checked at every assignment to `_TP` / `_PP` / `_ATTN_CP` / +`_DCP` / `_MOE_DP` in `parallel_state.py`. Three things do not follow that rule: + +- `initialize_model_parallel` aliases `_MOE_DP` to `_ATTN_CP` when `attn_cp_size > + moe_dp_size`, so a reader that means **the MoE communicator's width** calls + `get_moe_cp_size()`, not `get_parallel().moe_dp_size`. +- `patch_tensor_parallel_group` runs a scope under a different TP group (draft + workers), and declares it by overriding `tp_size`, `tp_rank` and `tp_group` + for the scope's duration. Readers inside need no special spelling. +- Elastic EP scales `ep_size` / `dp_size` on the published bag while the group + coordinators keep their construction width. Those are different names, not two + answers to one name. + +DCP keeps its own pair: `get_parallel().attn_dcp_size` / `.dcp_enabled` answer the +*effective* topology (`1` / `False` with no group installed), while `dcp_size` is +what the launch requested. A process-global seed field-read of one of these sizes (`get_server_args().tp_size`, or an alias of it) is a read-ratchet failure. A @@ -268,8 +275,7 @@ where an object was handed one; it is not a global accessor. - **a resolved leaf** → its namespace bag (`get_exec().moe.moe_runner_backend`, `get_schedule().chunked_prefill_size`, …). Bag-backed reads — a leaf directly, or - a bag-derived accessor below, including the `get_parallel().config` hop — are - what see post-publish overrides. Only the + a bag-derived accessor below — are what see post-publish overrides. Only the instance-derived accessors (the ones with no leaf to read) answer from the startup record and therefore do not. - **a leaf the caller names at runtime** (a readback reporting a list of fields) @@ -295,18 +301,14 @@ where an object was handed one; it is not a global accessor. property with no bag of its own. A new derived member gets an accessor here rather than call sites reaching for the record, and only when the bag-derived shape above cannot express it. -- **what was *configured*, where the bare name is the live value** - → `get_parallel().config.{tp,pp,moe_dp,attn_cp,dcp}_size`. It reads the parallel - bag's own leaf, so it answers with the resolved configuration and follows a - post-publish override. The DCP live pair (`get_parallel().attn_dcp_size` / - `.dcp_enabled`) is a different question again: it answers the effective topology - (`1` / `False` when no group is installed), never the requested size, and it does - not *need* dist init to answer. Every (file, size) pair is registered - with its reason in `test_global_config_read_ratchet.py` - (`_CONFIGURED_SIZE_CALL_SITES`), and that test fails if the code and the list - disagree — a new file, or a new size in a listed file, has to be added — so a new - site needs both an answer the live property cannot give and an entry saying what - it is. +- **a parallel size** → `get_parallel().{tp,pp,moe_dp,attn_cp,dcp}_size`, which is + the parallel bag's own leaf: it answers with the resolved configuration and + follows a post-publish override. Two questions are *not* that, and have their + own spelling: the width of the MoE communicator you are about to collectively + operate on is `get_moe_cp_size()` (the `_MOE_DP = _ATTN_CP` alias makes it + differ), and the effective DCP topology is `get_parallel().attn_dcp_size` / + `.dcp_enabled` (`1` / `False` when no group is installed), which does not need + dist init to answer. - **this runner's resolved value** → the runner (`prefill_attention_backend_str`, `kv_cache_dtype_str`, `draft_attention_backend`, `num_fused_shared_experts` on the model). @@ -534,18 +536,14 @@ ONE thread — do not design for TBO threads that don't exist. instance attribute, plus the `getattr(..., "field")` spelling of each; a name computed at runtime or indirection deeper than a local name copy is census-tool territory, per the test's docstring). The scanner matches `get_server_args` by its - literal name, and the same file *bans* `import ... as` renames of it so that - matching stays sound. Exempt by owner - module only (`runtime_context.py`, `server_args.py`, `arg_groups/`). The same file - carries `_CONFIGURED_SIZE_CALL_SITES`, the (file, size) map of every - `get_parallel().config.` reader with the reason the live property - cannot serve it — a new file or a new size in a listed file must be added there. Its - subject set is *derived* (property names ∩ `parallel` NS leaves), and it resolves - every spelling of the call itself — an aliased import, a module-qualified receiver - (including the whole dotted path an unaliased `import` binds), a local bound to either - hop — so neither a rename nor a new shadowed size escapes it. - `TestParallelConfigReadSpellings` in that file runs each spelling, because a spelling - the scanner cannot resolve drops the read instead of failing anything. + literal name — bare or module-qualified (`ctx.get_server_args()`) — and + `TestNoRenamedAccessorImports` in the same file *bans* `import ... as` renames of it, + which is what makes literal-name matching sound. Exempt by owner + module only (`runtime_context.py`, `server_args.py`, `arg_groups/`). Two classes, + no more: `TestGlobalConfigReadRatchet` holds the two baselines and + `TestNoRenamedAccessorImports` holds the ban. There is no configured-size registry + here any longer — `get_parallel()` has one spelling per name, so a size read is not a + choice between two answers and nothing needs registering. 6. **Module-state ratchet** (`test_module_state_ratchet.py`): `global` statements in the flag-owning layers are pinned by name. A new module-level runtime global belongs on a flags group / resources slot instead; migrating a pinned survivor must shrink the pin. @@ -579,9 +577,10 @@ Never module-skip a test "until the migration settles" — seed the context inst form** (attribute-source ints get automatic-dynamic after the first size change). Bools (≤2 values) are tolerable in any form — see `ForwardFlags._GRAPH_VISIBLE`. Config-bag leaves are real instance attributes for - exactly this reason, and the parallel config tier is read through the plain - `ParallelContext.config` property for the same reason (`__getattr__` is - error-only, and `object.__getattribute__` graph-breaks). Before moving such state, + exactly this reason. Parallel leaves are the exception that was measured rather + than assumed: they come through `ParallelContext.__getattr__`, which traces + under `torch.compile(fullgraph=True)` (`object.__getattribute__` is the form + that graph-breaks, and it is not on this path). Before moving such state, prove its readers sit outside compile coverage; a piecewise-prefill boot of a small model is the fast check (recompile storms show as `torch._dynamo hit config.recompile_limit` during the compile pass). diff --git a/python/sglang/benchmark/one_batch.py b/python/sglang/benchmark/one_batch.py index e522eacb9c19..bdd2d9b8553e 100644 --- a/python/sglang/benchmark/one_batch.py +++ b/python/sglang/benchmark/one_batch.py @@ -899,9 +899,12 @@ def latency_test( initialize_fp4_gemm_config() if get_bool_env_var("SGLANG_SET_CPU_AFFINITY"): - parallel = get_parallel().config + parallel = get_parallel() set_gpu_proc_affinity( - parallel.pp_size, parallel.tp_size, parallel.nnodes, tp_rank + parallel.pp_size, + parallel.tp_size, + parallel.nnodes, + tp_rank, ) # Configure the logger diff --git a/python/sglang/compile_deep_gemm.py b/python/sglang/compile_deep_gemm.py index d4e9fd67d2ba..65d809442dcc 100644 --- a/python/sglang/compile_deep_gemm.py +++ b/python/sglang/compile_deep_gemm.py @@ -81,7 +81,7 @@ async def warm_up_compile( ) generate_req_input.bootstrap_host = [FAKE_BOOTSTRAP_HOST] * dp_size generate_req_input.bootstrap_room = [ - i * (2**63 // dp_size) + (i % get_parallel().config.tp_size) + i * (2**63 // dp_size) + (i % get_parallel().tp_size) for i in range(dp_size) ] else: diff --git a/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py b/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py index cbce8735ebf4..eb642a83b9fd 100644 --- a/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py +++ b/python/sglang/multimodal_gen/runtime/distributed/parallel_state.py @@ -701,7 +701,16 @@ def model_parallel_is_initialized() -> bool: @contextmanager def use_tensor_parallel_group(tp_group: GroupCoordinator): - """Use one TP group consistently across diffusion and reused SRT modules.""" + """Use one TP group consistently across diffusion and reused SRT modules. + + The scope replaces the module globals that ``get_tp_group()`` and srt's + ``get_tp_group()`` / ``get_attention_tp_group()`` read, and — like srt's + ``patch_tensor_parallel_group`` — the three members the runtime context + answers with, so that a size read from the published bag cannot disagree + with a rank read from the swapped group. + """ + from sglang.srt.runtime_context import get_parallel + old_tp_group = get_tp_group() import sglang.srt.distributed.parallel_state as srt_parallel_state @@ -712,7 +721,12 @@ def use_tensor_parallel_group(tp_group: GroupCoordinator): srt_parallel_state._TP = tp_group srt_parallel_state._ATTN_TP = tp_group try: - yield + with get_parallel().override( + tp_size=tp_group.world_size, + tp_rank=tp_group.rank_in_group, + tp_group=tp_group, + ): + yield finally: _TP = old_tp_group srt_parallel_state._TP = old_srt_tp_group diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py index 14a2375633f7..dab6e2855389 100644 --- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py +++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py @@ -291,7 +291,12 @@ def init_device_and_model(self) -> None: from sglang.srt.server_args import ServerArgs as SrtServerArgs if get_context()._server_args is None: - publish(SrtServerArgs(model_path="dummy"), role="diffusion_gpu_worker") + # srt reads the size from the configuration and the rank from the + # live group, so the dummy carries the width just installed. + publish( + SrtServerArgs(model_path="dummy", tp_size=self.server_args.tp_size), + role="diffusion_gpu_worker", + ) # set proc title if model_parallel_is_initialized(): diff --git a/python/sglang/multimodal_gen/test/unit/test_component_accuracy_parallel_runtime.py b/python/sglang/multimodal_gen/test/unit/test_component_accuracy_parallel_runtime.py index 919b7bc9dc9f..a33f862b4f4a 100644 --- a/python/sglang/multimodal_gen/test/unit/test_component_accuracy_parallel_runtime.py +++ b/python/sglang/multimodal_gen/test/unit/test_component_accuracy_parallel_runtime.py @@ -17,10 +17,19 @@ initialize_parallel_runtime, ) from sglang.srt.distributed import parallel_state as srt_parallel_state +from sglang.srt.runtime_context import get_parallel _UTILS = "sglang.multimodal_gen.test.single_test_file.component_accuracy.utils" +def _tp_group(world_size: int = 1, rank_in_group: int = 0) -> SimpleNamespace: + """A TP group handle carrying the two members the scope declares to the + runtime context (`use_tensor_parallel_group` overrides `tp_size` / + `tp_rank` / `tp_group` for its duration); the scope otherwise only stores + the handle and compares it by identity.""" + return SimpleNamespace(world_size=world_size, rank_in_group=rank_in_group) + + def _server_args(*, ulysses_degree: int, ring_degree: int) -> SimpleNamespace: return SimpleNamespace( tp_size=1, @@ -162,7 +171,7 @@ def test_srt_tp_groups_follow_encoder_folding_context(): original_diffusion_tp_group = object() original_srt_tp_group = object() original_srt_attention_tp_group = object() - folding_tp_group = object() + folding_tp_group = _tp_group(world_size=2, rank_in_group=1) with ( patch.object(parallel_state, "_TP", original_diffusion_tp_group), @@ -177,6 +186,9 @@ def test_srt_tp_groups_follow_encoder_folding_context(): assert parallel_state._TP is folding_tp_group assert srt_parallel_state._TP is folding_tp_group assert srt_parallel_state._ATTN_TP is folding_tp_group + assert get_parallel().tp_size == 2 + assert get_parallel().tp_rank == 1 + assert get_parallel().tp_group is folding_tp_group assert parallel_state._TP is original_diffusion_tp_group assert srt_parallel_state._TP is original_srt_tp_group @@ -185,8 +197,8 @@ def test_srt_tp_groups_follow_encoder_folding_context(): def test_encoder_folding_context_is_nested_and_restores_each_group(): original_tp_group = object() - outer_tp_group = object() - inner_tp_group = object() + outer_tp_group = _tp_group(world_size=4, rank_in_group=3) + inner_tp_group = _tp_group(world_size=2, rank_in_group=1) with ( patch.object(parallel_state, "_TP", original_tp_group), @@ -194,14 +206,19 @@ def test_encoder_folding_context_is_nested_and_restores_each_group(): patch.object(srt_parallel_state, "_ATTN_TP", original_tp_group), ): with parallel_state.use_tensor_parallel_group(outer_tp_group): + assert get_parallel().tp_size == 4 with parallel_state.use_tensor_parallel_group(inner_tp_group): assert parallel_state._TP is inner_tp_group assert srt_parallel_state._TP is inner_tp_group assert srt_parallel_state._ATTN_TP is inner_tp_group + assert get_parallel().tp_size == 2 + assert get_parallel().tp_rank == 1 assert parallel_state._TP is outer_tp_group assert srt_parallel_state._TP is outer_tp_group assert srt_parallel_state._ATTN_TP is outer_tp_group + assert get_parallel().tp_size == 4 + assert get_parallel().tp_rank == 3 assert parallel_state._TP is original_tp_group assert srt_parallel_state._TP is original_tp_group diff --git a/python/sglang/srt/configs/zaya.py b/python/sglang/srt/configs/zaya.py index a07a7145e83b..3ff25afe317d 100644 --- a/python/sglang/srt/configs/zaya.py +++ b/python/sglang/srt/configs/zaya.py @@ -268,7 +268,7 @@ def mamba2_cache_params(self) -> Optional[Mamba2CacheParams]: try: tp_size = get_parallel().tp_size - except (AssertionError, RuntimeError): + except (AssertionError, RuntimeError, ValueError): tp_size = 1 in_out_ch_full = ( diff --git a/python/sglang/srt/debug_utils/dumper.py b/python/sglang/srt/debug_utils/dumper.py index e6edb3f1fbe0..6bdef9ebdaa3 100644 --- a/python/sglang/srt/debug_utils/dumper.py +++ b/python/sglang/srt/debug_utils/dumper.py @@ -1735,8 +1735,8 @@ def collect_parallel_info(self) -> dict: info["moe_tp_rank"] = parallel.moe_tp_rank info["moe_tp_size"] = parallel.moe_tp_size info["moe_dp_rank"] = parallel.moe_dp_rank - info["moe_dp_size"] = parallel.moe_dp_size - except (AttributeError, AssertionError): + info["moe_dp_size"] = self._dp_attn.get_moe_cp_size() + except (AttributeError, AssertionError, ValueError): info["distributed_error"] = True try: @@ -1748,7 +1748,7 @@ def collect_parallel_info(self) -> dict: info["attn_dp_size"] = self._dp_attn.get_attention_dp_size() info["attn_cp_rank"] = parallel.attn_cp_rank info["attn_cp_size"] = parallel.attn_cp_size - except (AttributeError, AssertionError): + except (AttributeError, AssertionError, ValueError): info["dp_attention_error"] = True return info diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index 1a578600e858..7d1ce18efaa7 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -186,7 +186,7 @@ def __init__( self.system_dp_rank = ( self.kv_args.system_dp_rank if self.kv_args.system_dp_rank else 0 ) - self.pp_size = get_parallel().config.pp_size + self.pp_size = get_parallel().pp_size self.pp_rank = self.kv_args.pp_rank self.local_ip = get_local_ip_auto() cp_sharded_prefill = self.attn_cp_size > 1 and ( diff --git a/python/sglang/srt/disaggregation/encoder/grpc_server.py b/python/sglang/srt/disaggregation/encoder/grpc_server.py index 58004e0eb918..ae2f5aac34bf 100644 --- a/python/sglang/srt/disaggregation/encoder/grpc_server.py +++ b/python/sglang/srt/disaggregation/encoder/grpc_server.py @@ -221,7 +221,7 @@ async def serve_grpc_encoder(server_args: ServerArgs): ).to_tcp() send_sockets: List[zmq.Socket] = [] - for rank in range(1, get_parallel().config.tp_size): + for rank in range(1, get_parallel().tp_size): schedule_path = f"ipc:///tmp/{ipc_path_prefix}_schedule_{rank}" send_sockets.append( get_zmq_socket(zmq_ctx, zmq.PUSH, schedule_path, bind=False) diff --git a/python/sglang/srt/disaggregation/encoder/receiver.py b/python/sglang/srt/disaggregation/encoder/receiver.py index ace0bcd8da19..53bcff1fc024 100644 --- a/python/sglang/srt/disaggregation/encoder/receiver.py +++ b/python/sglang/srt/disaggregation/encoder/receiver.py @@ -1734,7 +1734,7 @@ def __init__( self.host = get_local_ip_auto(get_serving().host) self.pp_rank = pp_rank self.tp_rank = tp_rank - self.tp_size = get_parallel().config.tp_size + self.tp_size = get_parallel().tp_size self.tp_group = tp_group self.nnodes = server_args.nnodes self.hostname = get_local_ip_auto() diff --git a/python/sglang/srt/disaggregation/encoder/runtime.py b/python/sglang/srt/disaggregation/encoder/runtime.py index 22005ede8876..b05d7298f6ae 100644 --- a/python/sglang/srt/disaggregation/encoder/runtime.py +++ b/python/sglang/srt/disaggregation/encoder/runtime.py @@ -1530,7 +1530,7 @@ def launch_local_runtime(server_args: ServerArgs) -> EncoderRuntime: send_sockets: List[zmq.Socket] = [] tp_processes: List[mp.Process] = [] - for rank in range(1, get_parallel().config.tp_size): + for rank in range(1, get_parallel().tp_size): schedule_path = f"ipc:///tmp/{ipc_path_prefix}_schedule_{rank}" send_sockets.append( get_zmq_socket(zmq_context, zmq.PUSH, schedule_path, bind=False) @@ -1570,10 +1570,10 @@ def launch_dp_runtime(server_args: ServerArgs) -> DPDispatcher: HTTP uses this entry point today. gRPC can reuse it later without importing HTTP application state or Uvicorn. """ - if get_parallel().dp_size <= 1 or get_parallel().config.tp_size != 1: + if get_parallel().dp_size <= 1 or get_parallel().tp_size != 1: raise ValueError( "Encoder DP mode requires --dp-size > 1 and --tp-size 1; got " - f"dp_size={get_parallel().dp_size}, tp_size={get_parallel().config.tp_size}." + f"dp_size={get_parallel().dp_size}, tp_size={get_parallel().tp_size}." ) dp_size = get_parallel().dp_size logger.info(f"Launching encoder in DP mode: dp_size={dp_size}") diff --git a/python/sglang/srt/disaggregation/encoder/server.py b/python/sglang/srt/disaggregation/encoder/server.py index 34b1386db723..9b63cf7e0f83 100644 --- a/python/sglang/srt/disaggregation/encoder/server.py +++ b/python/sglang/srt/disaggregation/encoder/server.py @@ -451,7 +451,7 @@ def __init__( this instance's value, not a config change, so it travels as an argument.""" assert_published(server_args, role="encoder") - logger.info(f"init MMEncoder {rank}/{get_parallel().config.tp_size}") + logger.info(f"init MMEncoder {rank}/{get_parallel().tp_size}") self.server_args = server_args configure_media_url_security( get_mm().allowed_media_domains, @@ -492,14 +492,12 @@ def __init__( init_distributed_environment( backend=get_default_distributed_backend(self.device), - world_size=get_parallel().config.tp_size, + world_size=get_parallel().tp_size, rank=rank, distributed_init_method=dist_init_method, local_rank=rank, ) - initialize_model_parallel( - tensor_model_parallel_size=get_parallel().config.tp_size - ) + initialize_model_parallel(tensor_model_parallel_size=get_parallel().tp_size) initialize_dp_attention(server_args, self.model_config) self.model = load_model( @@ -557,7 +555,7 @@ def __init__( ) self.mm_global_cache = EmbeddingCacheController( rank, - get_parallel().config.tp_size, + get_parallel().tp_size, embedding_store=embedding_store, hidden_dims=self._embedding_dims, tp_group=get_tp_group().cpu_group, @@ -1035,7 +1033,7 @@ async def _prepare_encode_context( ) def _broadcast_global_cache_mask(self, mask_tensor: torch.Tensor): - if get_parallel().config.tp_size > 1: + if get_parallel().tp_size > 1: torch.distributed.broadcast( mask_tensor, src=0, diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index f72fb25fc97f..fb45a22204a2 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -405,8 +405,7 @@ def __init__( ): super().__init__(args, disaggregation_mode, server_args, is_mla_backend) self.transfer_source_rank = ( - self.kv_args.pp_rank * get_parallel().config.tp_size - + self.kv_args.engine_rank + self.kv_args.pp_rank * get_parallel().tp_size + self.kv_args.engine_rank ) self.kv_args.kv_data_mem_kinds = _normalize_kv_mem_kinds( getattr(self.kv_args, "kv_data_mem_kinds", None), diff --git a/python/sglang/srt/distributed/parallel_state.py b/python/sglang/srt/distributed/parallel_state.py index 5aeb7a57311d..2e8c0bafd07c 100644 --- a/python/sglang/srt/distributed/parallel_state.py +++ b/python/sglang/srt/distributed/parallel_state.py @@ -2803,14 +2803,19 @@ def model_parallel_is_initialized(): @contextmanager def patch_tensor_parallel_group(tp_group: GroupCoordinator): - """Patch the tp group temporarily until this function ends. + """Run under a different tensor-parallel group until this scope ends. - This method is for draft workers of speculative decoding to run draft model - with different tp degree from that of target model workers. + This is for draft workers of speculative decoding, which run the draft model + at the target's attention-TP width rather than its global TP width. + + The scope replaces both the module global that ``get_tp_group()`` reads and + the three members the runtime context answers with. Args: tp_group (GroupCoordinator): the tp group coordinator """ + from sglang.srt.runtime_context import get_parallel + global _TP_STATE_PATCHED assert not _TP_STATE_PATCHED, "Should not call when it's already patched" @@ -2819,9 +2824,13 @@ def patch_tensor_parallel_group(tp_group: GroupCoordinator): global _TP _TP = tp_group try: - yield + with get_parallel().override( + tp_size=tp_group.world_size, + tp_rank=tp_group.rank_in_group, + tp_group=tp_group, + ): + yield finally: - # restore the original state _TP_STATE_PATCHED = False _TP = old_tp_group diff --git a/python/sglang/srt/elastic_ep/elastic_ep.py b/python/sglang/srt/elastic_ep/elastic_ep.py index 2568415187b8..2b1b79f3377b 100644 --- a/python/sglang/srt/elastic_ep/elastic_ep.py +++ b/python/sglang/srt/elastic_ep/elastic_ep.py @@ -127,7 +127,7 @@ def _init_joiner_state(cls, inst: ElasticEPState) -> None: if get_exec().moe.ep_join_mode == "scale": inst.effective_ep_size = ( - get_parallel().ep_join_rank_offset + get_parallel().config.tp_size + get_parallel().ep_join_rank_offset + get_parallel().tp_size ) inst.original_ep_size = ( get_parallel().elastic_ep_initial_size diff --git a/python/sglang/srt/elastic_ep/expert_backup_manager.py b/python/sglang/srt/elastic_ep/expert_backup_manager.py index 14e573317263..52fb0d4ec491 100644 --- a/python/sglang/srt/elastic_ep/expert_backup_manager.py +++ b/python/sglang/srt/elastic_ep/expert_backup_manager.py @@ -72,7 +72,7 @@ def __init__(self, server_args: ServerArgs, port_args: PortArgs): # losing the initial PUB message due to slow joiners. num_ready_clients = 0 - while num_ready_clients < get_parallel().config.tp_size: + while num_ready_clients < get_parallel().tp_size: sock_recv(self.recv_from_expert_backup_client) num_ready_clients += 1 diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index b24dc7dc7dd1..1c06b383c63a 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -683,7 +683,7 @@ def _launch_weight_cache_daemons(cls, server_args: ServerArgs): pp_rank_range, tp_rank_range, pp_size_per_node, tp_size_per_node = ( _calculate_rank_ranges( server_args.nnodes, - get_parallel().config.pp_size, + get_parallel().pp_size, tp_size, server_args.node_rank, ) @@ -844,7 +844,7 @@ def _launch_scheduler_processes( pp_rank_range, tp_rank_range, pp_size_per_node, tp_size_per_node = ( _calculate_rank_ranges( server_args.nnodes, - get_parallel().config.pp_size, + get_parallel().pp_size, server_args.tp_size, server_args.node_rank, ) @@ -1843,8 +1843,8 @@ def _compute_parallelism_ranks( """ attn_dp_size = get_parallel().dp_size if get_parallel().enable_dp_attention else 1 tp_size = server_args.tp_size - attn_cp_size = get_parallel().config.attn_cp_size - moe_dp_size = get_parallel().config.moe_dp_size + attn_cp_size = get_parallel().attn_cp_size + moe_dp_size = get_parallel().moe_dp_size # Parallelism hierarchy (outermost to innermost): # - Attention: Global(TP) -> DP -> ATTN_CP -> ATTN_TP (innermost) diff --git a/python/sglang/srt/entrypoints/v1_loads.py b/python/sglang/srt/entrypoints/v1_loads.py index a3efd49abad9..e3a6c27d298e 100644 --- a/python/sglang/srt/entrypoints/v1_loads.py +++ b/python/sglang/srt/entrypoints/v1_loads.py @@ -146,8 +146,8 @@ async def get_loads( "version": __version__, "accelerator": _accelerator_name(), "num_accelerators": _num_accelerators_per_dp_rank( - get_parallel().config.tp_size, - get_parallel().config.pp_size, + get_parallel().tp_size, + get_parallel().pp_size, get_parallel().dp_size, get_parallel().enable_dp_attention, ), diff --git a/python/sglang/srt/eplb/expert_location.py b/python/sglang/srt/eplb/expert_location.py index 4415dc348067..b2ced8192f07 100644 --- a/python/sglang/srt/eplb/expert_location.py +++ b/python/sglang/srt/eplb/expert_location.py @@ -245,7 +245,7 @@ def _init_common(model_config: ModelConfig): if get_exec().moe.ep_join_mode == "scale": ep_size = max( ep_size, - get_parallel().ep_join_rank_offset + get_parallel().config.tp_size, + get_parallel().ep_join_rank_offset + get_parallel().tp_size, ) num_physical_experts, num_local_physical_experts = ( _compute_elastic_expert_layout( diff --git a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py index 2dd615d00f3d..f077320659bf 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_indexer.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_indexer.py @@ -248,7 +248,7 @@ def __init__( if _is_cuda: self.sm_count = deep_gemm.get_num_sms() self.half_device_sm_count = ceil_align(self.sm_count // 2, 8) - pp_size = get_parallel().config.pp_size + pp_size = get_parallel().pp_size self.logits_with_pp_recv = pp_size > 1 and not get_pp_group().is_last_rank else: self.logits_with_pp_recv = False diff --git a/python/sglang/srt/layers/cp/base.py b/python/sglang/srt/layers/cp/base.py index 225fc4fdb399..308d06addc0e 100644 --- a/python/sglang/srt/layers/cp/base.py +++ b/python/sglang/srt/layers/cp/base.py @@ -282,15 +282,19 @@ def get_cp_strategy() -> Optional[ContextParallelStrategy]: global _STRATEGY if _STRATEGY is None: + # The reads are what raise, so they sit inside the guard. try: - parallel = get_parallel().config - except ValueError: + parallel = get_parallel() + enable_prefill_cp = parallel.enable_prefill_cp + cp_size = parallel.attn_cp_size + cp_strategy = parallel.cp_strategy + except (AssertionError, AttributeError, RuntimeError, ValueError): return None - if parallel.enable_prefill_cp: + if enable_prefill_cp: init_cp_strategy( enable_prefill_cp=True, - cp_size=parallel.attn_cp_size, - cp_strategy=parallel.cp_strategy, + cp_size=cp_size, + cp_strategy=cp_strategy, ) return _STRATEGY diff --git a/python/sglang/srt/layers/dp_attention.py b/python/sglang/srt/layers/dp_attention.py index ee739faef69b..2d7efc68396b 100644 --- a/python/sglang/srt/layers/dp_attention.py +++ b/python/sglang/srt/layers/dp_attention.py @@ -349,7 +349,7 @@ def initialize_dp_attention( ) enable_dp_attention = get_parallel().enable_dp_attention dp_size = get_parallel().dp_size - attn_cp_size = get_parallel().config.attn_cp_size + attn_cp_size = get_parallel().attn_cp_size dp.enabled = enable_dp_attention @@ -1026,12 +1026,10 @@ def get_moe_cp_size() -> int: def is_enable_moe_cp_allgather() -> bool: """True when moe_dp_size < attn_cp_size, requiring allgather across CP ranks before MoE. - Reads the configured sizes, not the live groups: that very configuration makes - ``initialize_model_parallel`` alias ``_MOE_DP`` to ``_ATTN_CP`` - (``parallel_state.py``), so the live sizes are equal and the comparison would - always be false. + In that configuration ``initialize_model_parallel`` aliases ``_MOE_DP`` to + ``_ATTN_CP``, so the two groups report equal widths. """ - return get_parallel().config.attn_cp_size > get_parallel().config.moe_dp_size + return get_parallel().attn_cp_size > get_parallel().moe_dp_size def moe_cp_all_gather_into_tensor(output: torch.Tensor, input: torch.Tensor): diff --git a/python/sglang/srt/layers/quantization/fp8_utils.py b/python/sglang/srt/layers/quantization/fp8_utils.py index 267bf3b976a6..bbd70a438d52 100755 --- a/python/sglang/srt/layers/quantization/fp8_utils.py +++ b/python/sglang/srt/layers/quantization/fp8_utils.py @@ -2120,7 +2120,9 @@ def validate_fp8_block_shape( ) -> None: """Validate block quantization shapes for tensor parallelism.""" - tp_size = getattr(layer, "tp_size", get_parallel().tp_size) + # Lazy: a ``getattr`` default would read the published bag even for a + # layer that carries its own tp_size. + tp_size = layer.tp_size if hasattr(layer, "tp_size") else get_parallel().tp_size block_n, block_k = block_size[0], block_size[1] # Required by row parallel diff --git a/python/sglang/srt/managers/data_parallel_controller.py b/python/sglang/srt/managers/data_parallel_controller.py index 13d795354fcc..b0d1320c1dc5 100644 --- a/python/sglang/srt/managers/data_parallel_controller.py +++ b/python/sglang/srt/managers/data_parallel_controller.py @@ -392,9 +392,7 @@ def launch_dp_schedulers(self, server_args, port_args): ) threads.append(thread) base_gpu_id += ( - server_args.tp_size - * get_parallel().config.pp_size - * server_args.gpu_id_step + server_args.tp_size * get_parallel().pp_size * server_args.gpu_id_step ) if server_args.node_rank == 0: @@ -615,8 +613,8 @@ def launch_tensor_parallel_group( scheduler_pipe_readers = [] - pp_size_per_node = max(get_parallel().config.pp_size // server_args.nnodes, 1) - nnodes_per_pp_rank = max(server_args.nnodes // get_parallel().config.pp_size, 1) + pp_size_per_node = max(get_parallel().pp_size // server_args.nnodes, 1) + nnodes_per_pp_rank = max(server_args.nnodes // get_parallel().pp_size, 1) pp_rank_range = range( pp_size_per_node * (server_args.node_rank // nnodes_per_pp_rank), pp_size_per_node * (server_args.node_rank // nnodes_per_pp_rank + 1), @@ -647,7 +645,7 @@ def launch_tensor_parallel_group( tp_rank, server_args.tp_size, get_parallel().dp_size, - get_parallel().config.attn_cp_size, + get_parallel().attn_cp_size, ) # compute zmq ports for this dp rank rank_port_args = PortArgs.init_new( @@ -683,22 +681,18 @@ def launch_tensor_parallel_group( # - Attention: Global(TP) -> DP -> ATTN_CP -> ATTN_TP (innermost) # - MoE: Global(TP) -> MOE_DP -> EP -> MOE_TP (innermost) attn_tp_size = ( - server_args.tp_size - // attn_dp_size - // get_parallel().config.attn_cp_size + server_args.tp_size // attn_dp_size // get_parallel().attn_cp_size ) - attn_cp_rank = ( - tp_rank // attn_tp_size - ) % get_parallel().config.attn_cp_size + attn_cp_rank = (tp_rank // attn_tp_size) % get_parallel().attn_cp_size moe_dp_rank = tp_rank // ( - server_args.tp_size // get_parallel().config.moe_dp_size + server_args.tp_size // get_parallel().moe_dp_size ) moe_ep_rank = ( tp_rank - % (server_args.tp_size // get_parallel().config.moe_dp_size) + % (server_args.tp_size // get_parallel().moe_dp_size) // ( server_args.tp_size - // get_parallel().config.moe_dp_size + // get_parallel().moe_dp_size // get_parallel().ep_size ) ) diff --git a/python/sglang/srt/managers/rust_server.py b/python/sglang/srt/managers/rust_server.py index 4251f15257b2..d97eb26d0337 100644 --- a/python/sglang/srt/managers/rust_server.py +++ b/python/sglang/srt/managers/rust_server.py @@ -275,7 +275,7 @@ def _use_feature_shm(self) -> bool: ) return ( - get_parallel().config.tp_size > 1 + get_parallel().tp_size > 1 and determine_tensor_transport_mode() != "default" and not self.server_args.skip_tokenizer_init ) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 9776f3a68625..29bd57bb1313 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -478,30 +478,30 @@ def __init__( compute_dp_attention_world_info( get_parallel().enable_dp_attention, tp_rank, - get_parallel().config.tp_size, + get_parallel().tp_size, get_parallel().dp_size, - get_parallel().config.attn_cp_size, + get_parallel().attn_cp_size, ) ) self.ps = ParallelState( tp_rank=tp_rank, - tp_size=get_parallel().config.tp_size, + tp_size=get_parallel().tp_size, pp_rank=pp_rank, - pp_size=get_parallel().config.pp_size, + pp_size=get_parallel().pp_size, dp_rank=dp_rank, dp_size=get_parallel().dp_size, attn_tp_rank=attn_tp_rank, attn_tp_size=attn_tp_size, attn_cp_rank=attn_cp_rank, - attn_cp_size=get_parallel().config.attn_cp_size, - attn_dcp_rank=tp_rank % get_parallel().config.dcp_size, - attn_dcp_size=get_parallel().config.dcp_size, + attn_cp_size=get_parallel().attn_cp_size, + attn_dcp_rank=tp_rank % get_parallel().dcp_size, + attn_dcp_size=get_parallel().dcp_size, attn_dp_rank=attn_dp_rank, attn_dp_size=attn_dp_size, moe_ep_rank=moe_ep_rank, moe_ep_size=get_parallel().ep_size, moe_dp_rank=moe_dp_rank, - moe_dp_size=get_parallel().config.moe_dp_size, + moe_dp_size=get_parallel().moe_dp_size, gpu_id=gpu_id, ) @@ -4544,7 +4544,12 @@ def get_internal_state(self, recv_req: GetInternalStateReq): # Resolved config (pristine server_args + post-publish overrides) so a # readback reflects values changed via /set_internal_state, not startup. ret = get_context().resolved_server_args_dict() - ret["world_size"] = compute_world_size(get_parallel().config) + ret["world_size"] = compute_world_size( + enable_dp_attention=get_parallel().enable_dp_attention, + dp_size=get_parallel().dp_size, + tp_size=get_parallel().tp_size, + pp_size=get_parallel().pp_size, + ) ret["last_gen_throughput"] = self.metrics_reporter.last_gen_throughput draft_graph_memory_usage = ( None if self.draft_worker is None else self.draft_worker.graph_memory_usage @@ -5220,7 +5225,7 @@ def dispatch_event_loop(scheduler: Scheduler): if disaggregation_mode == DisaggregationMode.NULL: if scheduler.enable_pdmux: scheduler.event_loop_pdmux() - elif get_parallel().config.pp_size > 1: + elif get_parallel().pp_size > 1: scheduler.event_loop_pp() elif scheduler.enable_overlap_mlx: scheduler.event_loop_overlap_mlx() @@ -5229,14 +5234,14 @@ def dispatch_event_loop(scheduler: Scheduler): else: scheduler.event_loop_normal() elif disaggregation_mode == DisaggregationMode.PREFILL: - if get_parallel().config.pp_size > 1: + if get_parallel().pp_size > 1: scheduler.event_loop_pp_disagg_prefill() elif scheduler.enable_overlap: scheduler.event_loop_overlap_disagg_prefill() else: scheduler.event_loop_normal_disagg_prefill() elif disaggregation_mode == DisaggregationMode.DECODE: - if get_parallel().config.pp_size > 1: + if get_parallel().pp_size > 1: scheduler.event_loop_pp_disagg_decode() elif scheduler.enable_overlap: scheduler.event_loop_overlap_disagg_decode() @@ -5277,13 +5282,13 @@ def configure_scheduler_process( prefix = "" if shown_dp is not None: prefix += f" DP{shown_dp}" - if get_parallel().config.pp_size > 1: + if get_parallel().pp_size > 1: prefix += f" PP{pp_rank}" - if get_parallel().config.attn_cp_size > 1: + if get_parallel().attn_cp_size > 1: prefix += f" ATTN_CP{attn_cp_rank}" - if get_parallel().config.moe_dp_size > 1: + if get_parallel().moe_dp_size > 1: prefix += f" MOE_DP{moe_dp_rank}" - if get_parallel().config.tp_size > 1: + if get_parallel().tp_size > 1: prefix += f" TP{shown_tp}" if get_parallel().ep_size > 1: prefix += f" EP{shown_moe_ep}" @@ -5299,8 +5304,8 @@ def configure_scheduler_process( # Set cpu affinity to this gpu process if envs.SGLANG_SET_CPU_AFFINITY.get(): set_gpu_proc_affinity( - get_parallel().config.pp_size, - get_parallel().config.tp_size, + get_parallel().pp_size, + get_parallel().tp_size, get_parallel().nnodes, gpu_id, ) diff --git a/python/sglang/srt/managers/scheduler_components/metrics_reporter.py b/python/sglang/srt/managers/scheduler_components/metrics_reporter.py index 1e535716253e..1d1f0ddf7c36 100644 --- a/python/sglang/srt/managers/scheduler_components/metrics_reporter.py +++ b/python/sglang/srt/managers/scheduler_components/metrics_reporter.py @@ -1117,7 +1117,7 @@ def _update_lora_metrics(self): active_lora_ids = set() # For PP mode, check all running micro batches - if get_parallel().config.pp_size > 1: + if get_parallel().pp_size > 1: for batch in self.scheduler.running_mbs: if batch and hasattr(batch, "reqs"): for req in batch.reqs: diff --git a/python/sglang/srt/managers/tokenizer_control_mixin.py b/python/sglang/srt/managers/tokenizer_control_mixin.py index 1e5cf202b784..f02298037c53 100644 --- a/python/sglang/srt/managers/tokenizer_control_mixin.py +++ b/python/sglang/srt/managers/tokenizer_control_mixin.py @@ -179,8 +179,8 @@ def update_control_communicator_fan_out(self: TokenizerManager, worker_count: in ) if primary_group_control: control_fan_out = ( - worker_count + get_parallel().config.tp_size - 1 - ) // get_parallel().config.tp_size + worker_count + get_parallel().tp_size - 1 + ) // get_parallel().tp_size else: control_fan_out = worker_count diff --git a/python/sglang/srt/mem_cache/kv_cache_configurator.py b/python/sglang/srt/mem_cache/kv_cache_configurator.py index 976fcf899d3f..3328fb28b598 100644 --- a/python/sglang/srt/mem_cache/kv_cache_configurator.py +++ b/python/sglang/srt/mem_cache/kv_cache_configurator.py @@ -1906,7 +1906,7 @@ def _apply_token_constraints(self, token_capacity: int) -> int: token_capacity = min(token_capacity, user_limit) # Sync across PP ranks (each may have different layer counts) - if get_parallel().config.pp_size > 1: + if get_parallel().pp_size > 1: tensor = torch.tensor(token_capacity, dtype=torch.int64) torch.distributed.all_reduce( tensor, diff --git a/python/sglang/srt/model_executor/cpu_graph_runner.py b/python/sglang/srt/model_executor/cpu_graph_runner.py index a8bc8aebf965..d2a0a72d7bcf 100644 --- a/python/sglang/srt/model_executor/cpu_graph_runner.py +++ b/python/sglang/srt/model_executor/cpu_graph_runner.py @@ -609,9 +609,9 @@ def __init__(self, model_runner: ModelRunner): self.enable_profile_cuda_graph = ( model_runner.server_args.enable_profile_cuda_graph ) - self.tp_size = get_parallel().config.tp_size + self.tp_size = get_parallel().tp_size self.dp_size = get_parallel().dp_size - self.pp_size = get_parallel().config.pp_size + self.pp_size = get_parallel().pp_size self.capture_forward_mode = ForwardMode.DECODE self.capture_hidden_mode = self.return_hidden_states_mode diff --git a/python/sglang/srt/model_executor/model_runner_components/spec_aux_hidden_state.py b/python/sglang/srt/model_executor/model_runner_components/spec_aux_hidden_state.py index 61a2a573b417..eb1028a4153e 100644 --- a/python/sglang/srt/model_executor/model_runner_components/spec_aux_hidden_state.py +++ b/python/sglang/srt/model_executor/model_runner_components/spec_aux_hidden_state.py @@ -240,7 +240,7 @@ def _resolve_dflash_draft_cell_size( draft_model_config=draft_model_config, draft_num_layers=draft_num_layers, draft_kv_cache_dtype=draft_kv_cache_dtype, - tp_size=get_parallel().config.tp_size, + tp_size=get_parallel().tp_size, ) except Exception as e: # noqa: BLE001 logger.warning( diff --git a/python/sglang/srt/model_executor/model_runner_components/startup_weight_load.py b/python/sglang/srt/model_executor/model_runner_components/startup_weight_load.py index 40899a4ba0ba..77c844d41837 100644 --- a/python/sglang/srt/model_executor/model_runner_components/startup_weight_load.py +++ b/python/sglang/srt/model_executor/model_runner_components/startup_weight_load.py @@ -119,10 +119,10 @@ def from_published_config( prefill_cuda_graph_backend=cuda_graph_config.prefill.backend, is_draft_worker=is_draft_worker, speculative_algorithm=get_spec().speculative_algorithm, - tp_size=get_parallel().config.tp_size, - attn_cp_size=get_parallel().config.attn_cp_size, - dcp_size=get_parallel().config.dcp_size, - pp_size=get_parallel().config.pp_size, + tp_size=get_parallel().tp_size, + attn_cp_size=get_parallel().attn_cp_size, + dcp_size=get_parallel().dcp_size, + pp_size=get_parallel().pp_size, dp_size=get_parallel().dp_size, ep_size=get_parallel().ep_size, cpu_offload_gb=get_exec().offload.cpu_offload_gb, diff --git a/python/sglang/srt/model_executor/runner/base_runner.py b/python/sglang/srt/model_executor/runner/base_runner.py index 3da934b24a39..15153a3f3bf4 100644 --- a/python/sglang/srt/model_executor/runner/base_runner.py +++ b/python/sglang/srt/model_executor/runner/base_runner.py @@ -216,10 +216,10 @@ def __init__(self, model_runner: ModelRunner) -> None: self.model_runner = model_runner self.device = model_runner.device self.device_module = torch.get_device_module(self.device) - self.tp_size = get_parallel().config.tp_size + self.tp_size = get_parallel().tp_size # elastic-EP scale-up rewrites dp_size on the published config self.dp_size = get_parallel().dp_size - self.pp_size = get_parallel().config.pp_size + self.pp_size = get_parallel().pp_size self.enable_pdmux = model_runner.server_args.enable_pdmux self.return_hidden_states_mode = ( CaptureHiddenMode.NULL @@ -349,7 +349,7 @@ def _alloc_dummy_decode_buffers( vocab_size=mr.model_config.vocab_size, dtype=mr.model_config.dtype, dp_size=get_parallel().dp_size, - pp_size=get_parallel().config.pp_size, + pp_size=get_parallel().pp_size, is_encoder_decoder=mr.model_config.is_encoder_decoder, require_mlp_tp_gather=require_mlp_tp_gather(), seq_len_fill_value=mr.attn_backend.get_cuda_graph_seq_len_fill_value(), @@ -521,7 +521,7 @@ def _dummy_run( extend_prefix_lens = None extend_start_loc = None - if get_parallel().config.pp_size > 1: + if get_parallel().pp_size > 1: # PP0 already cp-split hidden_states before send. pp_hidden_tokens = num_tokens if ( @@ -645,7 +645,7 @@ def run_once(): kwargs = {} if ( - get_parallel().config.pp_size > 1 + get_parallel().pp_size > 1 and "pp_proxy_tensors" in inspect.signature(mr.model.forward).parameters ): kwargs["pp_proxy_tensors"] = PPProxyTensors( diff --git a/python/sglang/srt/model_loader/loader.py b/python/sglang/srt/model_loader/loader.py index c5629381c8fb..613b823d626c 100644 --- a/python/sglang/srt/model_loader/loader.py +++ b/python/sglang/srt/model_loader/loader.py @@ -1884,17 +1884,19 @@ def _collect_shard_config(self, model_config: ModelConfig) -> Dict[str, Any]: def _safe(fn) -> int: try: return fn() - except (AssertionError, AttributeError, RuntimeError): + except (AssertionError, AttributeError, RuntimeError, ValueError): return 1 + from sglang.srt.layers.dp_attention import get_moe_cp_size + parallel = get_parallel() return { "tp": _safe(lambda: parallel.tp_size), - "dp": _safe(lambda: parallel.moe_dp_size), + "dp": _safe(get_moe_cp_size), "ep": _safe(lambda: parallel.moe_ep_size), "pp": _safe(lambda: parallel.pp_size), "moe_dense_tp_size": parallel.moe_dense_tp_size, - "moe_dp_size": get_parallel().config.moe_dp_size, + "moe_dp_size": get_parallel().moe_dp_size, "enable_dp_lm_head": parallel.enable_dp_lm_head, "enable_fp32_lm_head": get_exec().features.enable_fp32_lm_head, "quantization": model_config.quantization, diff --git a/python/sglang/srt/models/kimi_k25.py b/python/sglang/srt/models/kimi_k25.py index a98a1fd7cf2d..debbf587b754 100644 --- a/python/sglang/srt/models/kimi_k25.py +++ b/python/sglang/srt/models/kimi_k25.py @@ -733,7 +733,7 @@ def materialize_item_features(image_indices: List[int]) -> torch.Tensor: # Match the configured TP consumer count captured when the # tokenizer creates MmItemMemoryPool. A live attention subgroup # size could leave acknowledgements missing and strand the lease. - ipc_consumer_count = max(get_parallel().config.tp_size, 1) + ipc_consumer_count = max(get_parallel().tp_size, 1) device_index = device.index if device.type == "cuda" and device_index is None: device_index = torch.cuda.current_device() diff --git a/python/sglang/srt/models/kimi_k3.py b/python/sglang/srt/models/kimi_k3.py index f00ae8fcd65a..6e6d1a3e234c 100644 --- a/python/sglang/srt/models/kimi_k3.py +++ b/python/sglang/srt/models/kimi_k3.py @@ -3382,7 +3382,7 @@ def materialize_item_features(image_indices: List[int]) -> torch.Tensor: # Match the configured TP consumer count captured when the # tokenizer creates MmItemMemoryPool. A live attention subgroup # size could leave acknowledgements missing and strand the lease. - ipc_consumer_count = max(get_parallel().config.tp_size, 1) + ipc_consumer_count = max(get_parallel().tp_size, 1) device_index = device.index if device.type == "cuda" and device_index is None: device_index = torch.cuda.current_device() diff --git a/python/sglang/srt/ray/data_parallel_controller.py b/python/sglang/srt/ray/data_parallel_controller.py index 3b299e74eb03..e0528e191b5c 100644 --- a/python/sglang/srt/ray/data_parallel_controller.py +++ b/python/sglang/srt/ray/data_parallel_controller.py @@ -147,8 +147,8 @@ def _launch_ray_tp_group( bundle_idx = self.bundle_for_node[node_idx] pp_range, tp_range, pp_per_node, tp_per_node = _calculate_rank_ranges( nnodes, - get_parallel().config.pp_size, - get_parallel().config.tp_size, + get_parallel().pp_size, + get_parallel().tp_size, node_rank=node_idx, ) for pp_rank in pp_range: @@ -160,7 +160,7 @@ def _launch_ray_tp_group( tp_rank % tp_per_node ) - parallel = get_parallel().config + parallel = get_parallel() if parallel.enable_dp_attention: _, _, actual_dp_rank, _ = compute_dp_attention_world_info( parallel.enable_dp_attention, @@ -208,7 +208,7 @@ def _launch_ray_tp_group( world_size = _compute_world_size() bundle_indices = _resolve_bundle_indices(self.pg, world_size) - parallel = get_parallel().config + parallel = get_parallel() ranks_per_tp_group = parallel.tp_size * parallel.pp_size if dp_rank is not None: start_rank = dp_rank * ranks_per_tp_group @@ -237,9 +237,9 @@ def _launch_ray_tp_group( _, _, actual_dp_rank, _ = compute_dp_attention_world_info( get_parallel().enable_dp_attention, tp_rank, - get_parallel().config.tp_size, + get_parallel().tp_size, get_parallel().dp_size, - get_parallel().config.attn_cp_size, + get_parallel().attn_cp_size, ) rank_port_args = PortArgs.init_new( server_args, actual_dp_rank, worker_ports diff --git a/python/sglang/srt/ray/engine.py b/python/sglang/srt/ray/engine.py index 01e5e11ab8a6..5e875b589a68 100644 --- a/python/sglang/srt/ray/engine.py +++ b/python/sglang/srt/ray/engine.py @@ -111,7 +111,12 @@ def _compute_world_size() -> int: Reads the published parallel leaves: the driver is sizing the actors that will hold the process groups, so there is nothing live to ask. """ - return compute_world_size(get_parallel().config) + return compute_world_size( + enable_dp_attention=get_parallel().enable_dp_attention, + dp_size=get_parallel().dp_size, + tp_size=get_parallel().tp_size, + pp_size=get_parallel().pp_size, + ) def _resolve_bundle_indices(pg: PlacementGroup, world_size: int) -> List[int]: @@ -269,7 +274,7 @@ def _launch_scheduler_processes( placement_group as create_placement_group, ) - parallel = get_parallel().config + parallel = get_parallel() if parallel.enable_dp_attention: total_gpus = parallel.tp_size * parallel.pp_size else: @@ -332,8 +337,8 @@ def _launch_scheduler_processes( pp_range, tp_range, pp_per_node, tp_per_node = ( _calculate_rank_ranges( nnodes, - get_parallel().config.pp_size, - get_parallel().config.tp_size, + get_parallel().pp_size, + get_parallel().tp_size, node_rank=node_idx, ) ) @@ -369,7 +374,7 @@ def _launch_scheduler_processes( f"bundle_indices={bundle_indices}" ) - tp_size = get_parallel().config.tp_size + tp_size = get_parallel().tp_size for rank in range(world_size): pp_rank = rank // tp_size tp_rank = rank % tp_size @@ -448,7 +453,7 @@ def _launch_dp_scheduler_processes( RayDataParallelController, ) - parallel = get_parallel().config + parallel = get_parallel() if parallel.enable_dp_attention: # DP attention folds DP into TP — total GPUs = tp_size * pp_size total_gpus = parallel.tp_size * parallel.pp_size diff --git a/python/sglang/srt/runtime_context.py b/python/sglang/srt/runtime_context.py index 2ba19fa2ace1..5c0f85e61e41 100644 --- a/python/sglang/srt/runtime_context.py +++ b/python/sglang/srt/runtime_context.py @@ -13,16 +13,13 @@ # ============================================================================== """A single structured accessor for process-static runtime state. -``get_parallel()`` returns a ``ParallelContext`` whose bare attributes — tp / dcp -/ pp / moe / attn size and rank, plus the process-group handles — each delegate -live to the canonical getter in ``distributed.parallel_state`` / -``layers.dp_attention``. Returned values are exactly what those getters return; -this is a read-through wrapper, not a cache. It gives call-sites one import and -one naming scheme in place of a dozen free functions, plus a test-only -``override()`` hook to force a topology without monkeypatching the underlying -getters. The resolved parallel **configuration** is the same object's ``config`` -hop (``get_parallel().config.tp_size``), which reads the published ``parallel`` -bag: bare is the live group, ``config`` is what was configured. +``get_parallel()`` returns a ``ParallelContext``. Ranks and process-group handles +read through **live** to the canonical getter in ``distributed.parallel_state`` / +``layers.dp_attention`` — exactly what those getters return, a read-through +wrapper and not a cache. Every other name, the sizes included, is a leaf of the +published ``parallel`` bag. It gives call-sites one import and one naming scheme +in place of a dozen free functions, plus an ``override()`` hook to force a +topology without monkeypatching the underlying getters. ``get_server_args()`` returns the process-wide ``ServerArgs``. This is the user's raw input, kept **read-only** for debug and reproduction; what @@ -135,21 +132,26 @@ def _parallel_config_leaves() -> frozenset: class ParallelContext: - """Parallel-topology namespace: the live groups bare, configuration under - ``config``. - - ``get_parallel().tp_size`` and its size / rank / group siblings are - read-through ``@property`` over the canonical getters, so they answer with - the **live** process groups and raise before distributed init. The resolved - parallel **configuration** is one hop away, on the published bag: - ``get_parallel().config.tp_size``, ``.config.nccl_port``. It answers in any - process at any point after publish, and follows a post-publish ``override``. - - The two disagree by design, so which one a call site wants is spelled at the - call site — no ``config`` means live. Elastic EP scales the live world away - from the configured one, and ``initialize_model_parallel`` aliases - ``_MOE_DP`` to ``_ATTN_CP`` when ``attn_cp_size > moe_dp_size``, which makes - a live comparison of that pair degenerate. + """Parallel-topology namespace: one spelling per name. + + Ranks and group handles are read-through ``@property`` over the canonical + getters, so they answer with the **live** process groups and raise before + distributed init. Every other name — ``tp_size`` and its size siblings + included, alongside config-only leaves such as ``nccl_port`` — is answered + from the published ``parallel`` bag, in any process at any point after + publish. + + A size is read from the configuration because the groups are built at + exactly the configured widths. Two things do not follow that rule and are + asked of the group itself: ``initialize_model_parallel`` aliases ``_MOE_DP`` + to ``_ATTN_CP`` when ``attn_cp_size > moe_dp_size``, so a reader that means + the MoE communicator's width calls ``get_moe_cp_size()``; and + ``patch_tensor_parallel_group`` runs a scope under a different TP group, + which it declares by overriding ``tp_size``, ``tp_rank`` and ``tp_group`` + for its duration. Elastic EP is a third case, and it needs no rule here: it + scales ``ep_size`` / ``dp_size`` on the published bag while the group + coordinators keep the width they were constructed with, so the two are + different names rather than two answers to one name. """ __slots__ = ("_overrides", "_config") @@ -158,28 +160,15 @@ def __init__(self): self._overrides = {} self._config = None # parallel config bag, wired at publish - @property - def config(self) -> _ConfigBag: - """The published ``parallel`` config bag. - - Reads the slot directly: ``parallel`` sits outside the per-role - namespace table (every process reads topology config), so no role check - applies here. The body stays - dynamo-traceable — ``get_parallel().config.moe_dense_tp_size`` and the - gate helpers over it run inside compiled model forwards. - """ - config = self._config - if config is None: - raise ValueError("config namespace 'parallel' not published") - return config - def __getattr__(self, name): - # Reached only for names with no live @property: the bare config leaves. if name.startswith("_"): # This also breaks the recursion when the ``_config`` slot itself is # still unset (pickle/copy protocols probe attributes before # __init__ runs). raise AttributeError(name) + overrides = self._overrides + if name in overrides: + return overrides[name] config = self._config if config is not None: if name in config._fields: @@ -214,18 +203,10 @@ def world_size(self) -> int: def world_rank(self) -> int: return self._v("world_rank", _ps().get_world_rank) - @property - def tp_size(self) -> int: - return self._v("tp_size", _ps().get_tensor_model_parallel_world_size) - @property def tp_rank(self) -> int: return self._v("tp_rank", _ps().get_tensor_model_parallel_rank) - @property - def pp_size(self) -> int: - return self._v("pp_size", _ps().get_pipeline_model_parallel_world_size) - @property def pp_rank(self) -> int: return self._v("pp_rank", _ps().get_pipeline_model_parallel_rank) @@ -238,10 +219,6 @@ def moe_ep_size(self) -> int: def moe_ep_rank(self) -> int: return self._v("moe_ep_rank", _ps().get_moe_expert_parallel_rank) - @property - def moe_dp_size(self) -> int: - return self._v("moe_dp_size", _ps().get_moe_data_parallel_world_size) - @property def moe_dp_rank(self) -> int: return self._v("moe_dp_rank", _ps().get_moe_data_parallel_rank) @@ -262,18 +239,10 @@ def attn_tp_size(self) -> int: def attn_tp_rank(self) -> int: return self._v("attn_tp_rank", _ps().get_attn_tensor_model_parallel_rank) - @property - def attn_cp_size(self) -> int: - return self._v("attn_cp_size", _ps().get_attn_context_model_parallel_world_size) - @property def attn_cp_rank(self) -> int: return self._v("attn_cp_rank", _ps().get_attn_context_model_parallel_rank) - @property - def dcp_size(self) -> int: - return self._v("dcp_size", _ps().get_dcp_world_size) - @property def dcp_rank(self) -> int: return self._v("dcp_rank", _ps().get_dcp_rank) @@ -283,14 +252,15 @@ def dcp_enabled(self) -> bool: def getter(): if _ps().get_dcp_group_no_assert() is None: return False - return self.dcp_size > 1 + return _ps().get_dcp_world_size() > 1 return self._v("dcp_enabled", getter) @property def attn_dcp_size(self) -> int: return self._v( - "attn_dcp_size", lambda: self.dcp_size if self.dcp_enabled else 1 + "attn_dcp_size", + lambda: _ps().get_dcp_world_size() if self.dcp_enabled else 1, ) @property @@ -1163,8 +1133,8 @@ def get_forward() -> ForwardFlags: # --- Resolved config namespaces ------------------------- # Each returns the top-level snapshot bag; reads are `get_exec().moe.field` etc. # All fail with ValueError("... not published") until publish has projected them. -# ``parallel`` has no getter of its own: its bag is reached as -# ``get_parallel().config``, alongside the live topology it belongs to. +# ``parallel`` has no bag getter: ``get_parallel()`` answers its leaves +# directly, alongside the live topology they belong to. def get_device() -> _ConfigBag: return _CONTEXT.config_bag("device") @@ -1215,7 +1185,7 @@ def get_observability() -> _ConfigBag: # table declares which top-level config namespaces each role reads. ``None`` # means the full tree — either the role genuinely needs everything (scheduler) # or its deployment shape has not been audited yet (restrict only what smoke -# coverage can verify). ``parallel`` is served by ``get_parallel().config`` and +# coverage can verify). ``parallel`` is served by ``get_parallel()`` and # every process legitimately reads topology config, so it is not in this table. # # ``SGLANG_ROLE_NAMESPACES`` selects the mode (read once at import): @@ -1611,11 +1581,7 @@ def max_prefill_buffer_tokens() -> int: else 0 ) tokens = chunked - if ( - schedule.enable_dynamic_chunking - and get_parallel().config.pp_size > 1 - and chunked - ): + if schedule.enable_dynamic_chunking and get_parallel().pp_size > 1 and chunked: tokens = max( tokens, schedule.max_prefill_tokens or 0, math.ceil(chunked * 1.25) ) @@ -1645,7 +1611,7 @@ def pre_capture_activation_reserve_mb(gpu_mem: float | None) -> float: activation_tokens = max(schedule.chunked_prefill_size, 2048) else: activation_tokens = max(schedule.max_prefill_tokens, 2048) - parallel = get_parallel().config + parallel = get_parallel() reserved_mem = ( 512 + activation_tokens * 1.5 + parallel.tp_size * parallel.pp_size / 8 * 1024 ) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index dae349f7b587..22bdae06bcf8 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -11024,20 +11024,16 @@ def resolve_encoder_transfer_backend( return "zmq_to_scheduler" -def compute_world_size(config) -> int: - """Return the total GPU count across all data-parallel replicas. - - Takes the resolved topology -- the published `parallel` bag, or a view over - the declarations. `enable_dp_attention` and `dp_size` are both resolution's - answers (`_handle_dwdp` fills the pair, DeepSeek MLA context parallelism - turns DP attention on), so a raw-record read would size the world from what - the operator typed. +def compute_world_size( + *, enable_dp_attention: bool, dp_size: int, tp_size: int, pp_size: int +) -> int: + """Total GPU count across all data-parallel replicas. + + Takes the values rather than a config object: the two sizes are the widths + the launch asked for, which the Ray driver needs before any process group + exists, and passing a context would hand it the live groups instead. """ - return ( - (1 if config.enable_dp_attention else config.dp_size) - * config.tp_size - * config.pp_size - ) + return (1 if enable_dp_attention else dp_size) * tp_size * pp_size def m3_fp8_attn_gemm_enabled(args) -> bool: diff --git a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py index 8e89a86ce4ce..1ca19cbf6f0a 100644 --- a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py @@ -113,7 +113,7 @@ def __init__( self.device_module = torch.get_device_module(self.device) self.tp_size = model_runner.ps.tp_size self.attn_dp_size = model_runner.ps.attn_dp_size - self.pp_size = get_parallel().config.pp_size + self.pp_size = get_parallel().pp_size self.enable_torch_compile = get_flags().capture.enable_torch_compile self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.require_gathered_buffer = require_gathered_buffer() diff --git a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py index 2ee2f7c2805d..919b78aa4146 100644 --- a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py @@ -99,7 +99,7 @@ def __init__( self.device_module = torch.get_device_module(self.device) self.tp_size = model_runner.ps.tp_size self.attn_dp_size = model_runner.ps.attn_dp_size - self.pp_size = get_parallel().config.pp_size + self.pp_size = get_parallel().pp_size self.enable_torch_compile = get_flags().capture.enable_torch_compile self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.require_gathered_buffer = require_gathered_buffer() diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py b/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py index 43914af2a92f..3c3c5607ae88 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py @@ -99,7 +99,7 @@ def __init__(self, frozen_kv_mtp_worker: FrozenKVMTPDraftWorker): self.require_attn_tp_gather = require_attn_tp_gather() self.tp_size = self.model_runner.ps.tp_size self.attn_dp_size = self.model_runner.ps.attn_dp_size - self.pp_size = get_parallel().config.pp_size + self.pp_size = get_parallel().pp_size self.speculative_num_steps = get_spec().speculative_num_steps self.topk = get_spec().speculative_eagle_topk self.draft_attn_backend = frozen_kv_mtp_worker.draft_attn_backend diff --git a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py index 994686a2ed06..0579f95b82a0 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py @@ -155,7 +155,7 @@ def __init__(self, eagle_worker: MultiLayerEagleDraftWorker, step: int): self.device_module = torch.get_device_module(self.device) self.tp_size = model_runner.ps.tp_size self.dp_size = get_parallel().dp_size - self.pp_size = get_parallel().config.pp_size + self.pp_size = get_parallel().pp_size self.enable_torch_compile = get_flags().capture.enable_torch_compile self.disable_padding = model_runner.server_args.disable_cuda_graph_padding self.require_gathered_buffer = require_gathered_buffer() diff --git a/python/sglang/srt/utils/common.py b/python/sglang/srt/utils/common.py index 9e951adf3b89..28058bd26bf3 100644 --- a/python/sglang/srt/utils/common.py +++ b/python/sglang/srt/utils/common.py @@ -3756,7 +3756,7 @@ def require_mlp_tp_gather(): else: return ( get_parallel().moe_dense_tp_size - > get_parallel().config.tp_size // get_parallel().dp_size + > get_parallel().tp_size // get_parallel().dp_size ) else: return False @@ -3782,7 +3782,7 @@ def require_attn_tp_gather(): or get_parallel().moe_dense_tp_size is not None ): if get_parallel().enable_dp_attention: - return get_parallel().dp_size < get_parallel().config.tp_size + return get_parallel().dp_size < get_parallel().tp_size else: return True else: diff --git a/python/sglang/srt/utils/cuda_vmm_transport_utils.py b/python/sglang/srt/utils/cuda_vmm_transport_utils.py index fd79430d67d1..399f145a76da 100644 --- a/python/sglang/srt/utils/cuda_vmm_transport_utils.py +++ b/python/sglang/srt/utils/cuda_vmm_transport_utils.py @@ -163,8 +163,8 @@ def _contains_tensor_container(value) -> bool: def get_vmm_feature_consumer_count() -> int: if get_parallel().enable_dp_attention: - return get_parallel().config.tp_size // get_parallel().dp_size - return get_parallel().config.tp_size + return get_parallel().tp_size // get_parallel().dp_size + return get_parallel().tp_size class CudaVmmMemoryPool: diff --git a/python/sglang/srt/weight_cache/ipc_loader.py b/python/sglang/srt/weight_cache/ipc_loader.py index 2df1b5ad2b6a..4b7e90897f1c 100644 --- a/python/sglang/srt/weight_cache/ipc_loader.py +++ b/python/sglang/srt/weight_cache/ipc_loader.py @@ -494,6 +494,7 @@ def _fetch_from_cache(self, model_config) -> Optional[dict]: try: # Build engine's config fingerprint + from sglang.srt.layers.dp_attention import get_moe_cp_size from sglang.srt.runtime_context import get_exec, get_parallel ps = get_parallel() @@ -504,7 +505,7 @@ def _fetch_from_cache(self, model_config) -> Optional[dict]: pp_rank = ps.pp_rank ep_size = ps.moe_ep_size - moe_dp_size = ps.moe_dp_size + moe_dp_size = get_moe_cp_size() moe_dp_rank = ps.moe_dp_rank moe_ep_rank = ps.moe_ep_rank diff --git a/test/registered/layers/mamba/test_mamba2_mixer.py b/test/registered/layers/mamba/test_mamba2_mixer.py index 4cd95c6dd9a6..4a127d1b5e75 100644 --- a/test/registered/layers/mamba/test_mamba2_mixer.py +++ b/test/registered/layers/mamba/test_mamba2_mixer.py @@ -109,9 +109,12 @@ def mixer2_gated_norm_tensor_parallel( import sglang.srt.layers.attention.mamba.mixer2_rms_norm_gated as m2 - # Force attn-TP rank through the context (the weight loader reads it via - # get_parallel().attn_tp_rank); avoids calling initialize_dp_attention. - with get_parallel().override(attn_tp_rank=local_rank): + # Force the TP topology through the context (the weight loader reads + # get_parallel().attn_tp_rank, Mixer2RMSNormGated reads tp_size / tp_rank); + # avoids calling initialize_dp_attention. + with get_parallel().override( + attn_tp_rank=local_rank, tp_size=world_size, tp_rank=local_rank + ): # create gated-norm with TP mixer = m2.Mixer2RMSNormGated( full_hidden_size=hidden_size, diff --git a/test/registered/unit/layers/quantization/test_int8_linear_methods.py b/test/registered/unit/layers/quantization/test_int8_linear_methods.py index 57dbe83772bf..f6f3dd054ee9 100644 --- a/test/registered/unit/layers/quantization/test_int8_linear_methods.py +++ b/test/registered/unit/layers/quantization/test_int8_linear_methods.py @@ -12,6 +12,7 @@ from sglang.srt.layers.quantization.blockwise_int8 import BlockInt8Config from sglang.srt.layers.quantization.w8a8_int8 import W8A8Int8Config +from sglang.srt.runtime_context import get_parallel from sglang.srt.utils import get_device_sm from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.layer_ut_utils import ( @@ -112,7 +113,9 @@ def _build_layer(n: int, k: int): activation_scheme="dynamic", weight_block_size=[128, 128], ) - layer = make_tp1_column_parallel_linear(quant_config, n, k) + # create_weights reads get_parallel().tp_size, not the layer's argument. + with get_parallel().override(tp_size=1, tp_rank=0): + layer = make_tp1_column_parallel_linear(quant_config, n, k) w = torch.randn((n, k), device="cuda", dtype=torch.bfloat16) / 10 w_int8, scale_inv, w_dequant = _quantize_int8_block(w) load_linear_weights(layer, weight=w_int8, weight_scale_inv=scale_inv) diff --git a/test/registered/unit/managers/test_scheduler_internal_state_env_vars.py b/test/registered/unit/managers/test_scheduler_internal_state_env_vars.py index 1d6ff3276397..22c9f6c0e492 100644 --- a/test/registered/unit/managers/test_scheduler_internal_state_env_vars.py +++ b/test/registered/unit/managers/test_scheduler_internal_state_env_vars.py @@ -11,6 +11,7 @@ from sglang.srt.environ import envs from sglang.srt.managers.io_struct import GetInternalStateReq from sglang.srt.managers.scheduler import Scheduler +from sglang.srt.runtime_context import get_context register_cpu_ci(est_time=5, suite="base-a-test-cpu") @@ -42,18 +43,7 @@ def _get_internal_state(self) -> dict: ) scheduler.draft_worker = None - with patch( - "sglang.srt.managers.scheduler.get_context", - return_value=SimpleNamespace(resolved_server_args_dict=dict), - ), patch( - "sglang.srt.managers.scheduler.get_exec", - return_value=SimpleNamespace(moe=SimpleNamespace(elastic_ep_backend=None)), - ), patch( - "sglang.srt.managers.scheduler.compute_world_size", return_value=1 - ), patch( - "sglang.srt.managers.scheduler.get_parallel", - return_value=SimpleNamespace(config=SimpleNamespace()), - ): + with get_context().override_server_args(): output = scheduler.get_internal_state(recv_req=GetInternalStateReq()) return output.internal_state diff --git a/test/registered/unit/managers/test_scheduler_internal_state_world_size.py b/test/registered/unit/managers/test_scheduler_internal_state_world_size.py index a5141e29dc21..2b6dc1d8f889 100644 --- a/test/registered/unit/managers/test_scheduler_internal_state_world_size.py +++ b/test/registered/unit/managers/test_scheduler_internal_state_world_size.py @@ -1,6 +1,5 @@ import unittest from types import SimpleNamespace -from unittest.mock import patch from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import maybe_stub_sgl_kernel @@ -9,59 +8,52 @@ from sglang.srt.managers.io_struct import GetInternalStateReq from sglang.srt.managers.scheduler import Scheduler +from sglang.srt.runtime_context import get_context from sglang.srt.server_args import compute_world_size register_cpu_ci(est_time=5, suite="base-a-test-cpu") -def _make_parallel_config( +def _shape( *, tp_size: int, pp_size: int, dp_size: int, enable_dp_attention: bool -) -> SimpleNamespace: +) -> dict: """The four `parallel` leaves the world size is computed from.""" - return SimpleNamespace( - tp_size=tp_size, - pp_size=pp_size, - dp_size=dp_size, - enable_dp_attention=enable_dp_attention, - ) + return { + "tp_size": tp_size, + "pp_size": pp_size, + "dp_size": dp_size, + "enable_dp_attention": enable_dp_attention, + } class TestComputeWorldSize(unittest.TestCase): def test_a_single_gpu_server_holds_one_gpu(self): """The default shape has to come out as one, or every consumer is off by a factor.""" - config = _make_parallel_config( - tp_size=1, pp_size=1, dp_size=1, enable_dp_attention=False - ) + shape = _shape(tp_size=1, pp_size=1, dp_size=1, enable_dp_attention=False) - self.assertEqual(compute_world_size(config), 1) + self.assertEqual(compute_world_size(**shape), 1) def test_tensor_and_pipeline_stages_multiply(self): """Each (pp_rank, tp_rank) pair is its own scheduler process on its own gpu.""" - config = _make_parallel_config( - tp_size=2, pp_size=3, dp_size=1, enable_dp_attention=False - ) + shape = _shape(tp_size=2, pp_size=3, dp_size=1, enable_dp_attention=False) - self.assertEqual(compute_world_size(config), 6) + self.assertEqual(compute_world_size(**shape), 6) def test_plain_data_parallel_replicas_each_hold_their_own_gpus(self): """Without dp attention every replica launches a full tensor-parallel group of its own.""" - config = _make_parallel_config( - tp_size=2, pp_size=1, dp_size=2, enable_dp_attention=False - ) + shape = _shape(tp_size=2, pp_size=1, dp_size=2, enable_dp_attention=False) - self.assertEqual(compute_world_size(config), 4) + self.assertEqual(compute_world_size(**shape), 4) def test_data_parallel_attention_shares_the_tensor_parallel_gpus(self): """With dp attention the dp ranks live inside the tensor-parallel world, not beside it.""" - config = _make_parallel_config( - tp_size=4, pp_size=1, dp_size=2, enable_dp_attention=True - ) + shape = _shape(tp_size=4, pp_size=1, dp_size=2, enable_dp_attention=True) - self.assertEqual(compute_world_size(config), 4) + self.assertEqual(compute_world_size(**shape), 4) class TestSchedulerInternalStateWorldSize(unittest.TestCase): - def _get_internal_state(self, config: SimpleNamespace) -> dict: + def _get_internal_state(self, shape: dict) -> dict: scheduler = Scheduler.__new__(Scheduler) scheduler.metrics_reporter = SimpleNamespace( last_gen_throughput=1.0, @@ -87,40 +79,27 @@ def _get_internal_state(self, config: SimpleNamespace) -> dict: ) scheduler.draft_worker = None - with patch( - "sglang.srt.managers.scheduler.get_context", - return_value=SimpleNamespace(resolved_server_args_dict=dict), - ), patch( - "sglang.srt.managers.scheduler.get_exec", - return_value=SimpleNamespace(moe=SimpleNamespace(elastic_ep_backend=None)), - ), patch( - "sglang.srt.managers.scheduler.get_parallel", - return_value=SimpleNamespace(config=config), - ): + with get_context().override_server_args(**shape): output = scheduler.get_internal_state(recv_req=GetInternalStateReq()) return output.internal_state def test_the_internal_state_reports_the_whole_server(self): """A consumer sizing an external fleet reads the gpus the server occupies, not the declared sizes.""" - config = _make_parallel_config( - tp_size=2, pp_size=1, dp_size=2, enable_dp_attention=False - ) + shape = _shape(tp_size=2, pp_size=1, dp_size=2, enable_dp_attention=False) - internal_state = self._get_internal_state(config) + internal_state = self._get_internal_state(shape) self.assertEqual(internal_state["world_size"], 4) def test_the_reported_size_is_not_one_replica_of_a_data_parallel_server(self): """Each plain dp replica has its own process group, so no scheduler can report the whole server from it.""" - config = _make_parallel_config( - tp_size=2, pp_size=1, dp_size=2, enable_dp_attention=False - ) + shape = _shape(tp_size=2, pp_size=1, dp_size=2, enable_dp_attention=False) - internal_state = self._get_internal_state(config) + internal_state = self._get_internal_state(shape) self.assertNotEqual( - internal_state["world_size"], config.tp_size * config.pp_size + internal_state["world_size"], shape["tp_size"] * shape["pp_size"] ) 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 ad325bd97d4e..328f7e7fdfbf 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 @@ -63,79 +63,85 @@ def _run(rank: int, world: int, port: int): LayerSplitDSATokenToKVPool, ) - cp_rank = get_parallel().attn_cp_rank - cp_size = get_parallel().attn_cp_size - assert cp_size == world - - pool = LayerSplitDSATokenToKVPool( - SIZE, - page_size=PAGE_SIZE, - kv_lora_rank=KV_LORA_RANK, - dtype=torch.bfloat16, - qk_rope_head_dim=QK_ROPE, - layer_num=LAYER_NUM, - device=f"cuda:{rank}", - index_head_dim=INDEX_HEAD_DIM, - enable_memory_saver=False, - kv_cache_dim=KV_LORA_RANK + QK_ROPE, - layer_shard_rank=cp_rank, - layer_shard_size=cp_size, - ) - - # Owner writes a layer-distinct constant into each owned kv_buffer layer. - for layer_id in range(LAYER_NUM): - if pool._is_layer_owned(layer_id): - pool.kv_buffer[layer_id].fill_(float(layer_id + 1)) - - torch.cuda.synchronize() - torch.distributed.barrier() - - # Every rank reads every layer; broadcast must surface the owner's value. - ok = True - for layer_id in range(LAYER_NUM): - buf = pool._get_broadcastable_kv_buffer(layer_id) - expected = float(layer_id + 1) - got = buf.float().mean().item() - if abs(got - expected) > 1e-3: - print(f"[rank {rank}] layer {layer_id}: expected {expected}, got {got}") - ok = False - assert ok, f"rank {rank} read stale/incorrect broadcast contents" - - # Indexer buffer owner-broadcast: owner writes a layer-distinct value, then - # every rank must read it back for every layer. - for layer_id in range(LAYER_NUM): - store_buf = pool.get_index_k_with_scale_buffer(layer_id) - assert ( - store_buf.data_ptr() == pool.index_k_with_scale_buffer[layer_id].data_ptr() + # This worker builds the groups but never publishes, so the scope declares + # the size it runs at. + with get_parallel().override(attn_cp_size=world): + cp_rank = get_parallel().attn_cp_rank + cp_size = get_parallel().attn_cp_size + assert cp_size == world + + pool = LayerSplitDSATokenToKVPool( + SIZE, + page_size=PAGE_SIZE, + kv_lora_rank=KV_LORA_RANK, + dtype=torch.bfloat16, + qk_rope_head_dim=QK_ROPE, + layer_num=LAYER_NUM, + device=f"cuda:{rank}", + index_head_dim=INDEX_HEAD_DIM, + enable_memory_saver=False, + kv_cache_dim=KV_LORA_RANK + QK_ROPE, + layer_shard_rank=cp_rank, + layer_shard_size=cp_size, ) - if pool._is_layer_owned(layer_id): - store_buf.fill_(layer_id + 10) - torch.cuda.synchronize() - torch.distributed.barrier() - for layer_id in range(LAYER_NUM): - # invalidate any cached remote copy so the read forces a fresh broadcast - pool.invalidate_index_buffer_for_layer(layer_id) - buf = pool._get_broadcastable_index_buffer(layer_id) - expected = layer_id + 10 - got = buf.float().mean().item() - if abs(got - expected) > 1e-3: - print(f"[rank {rank}] index layer {layer_id}: exp {expected}, got {got}") - ok = False - assert ok, f"rank {rank} read stale/incorrect index broadcast contents" - - # Async prefetch path: prefetch layer, then read must return owner value. - for layer_id in range(LAYER_NUM): - pool.remote_kv_layer_id = None # force a fresh broadcast - pool.prefetch_kv_buffer(layer_id) - buf = pool._get_broadcastable_kv_buffer(layer_id) - got = buf.float().mean().item() - if abs(got - float(layer_id + 1)) > 1e-3: - print(f"[rank {rank}] prefetch layer {layer_id}: got {got}") - ok = False - assert ok, f"rank {rank} prefetch path returned incorrect contents" - - print(f"[rank {rank}] OK: all {LAYER_NUM} layers read correct owner contents") - torch.distributed.barrier() + + # Owner writes a layer-distinct constant into each owned kv_buffer layer. + for layer_id in range(LAYER_NUM): + if pool._is_layer_owned(layer_id): + pool.kv_buffer[layer_id].fill_(float(layer_id + 1)) + + torch.cuda.synchronize() + torch.distributed.barrier() + + # Every rank reads every layer; broadcast must surface the owner's value. + ok = True + for layer_id in range(LAYER_NUM): + buf = pool._get_broadcastable_kv_buffer(layer_id) + expected = float(layer_id + 1) + got = buf.float().mean().item() + if abs(got - expected) > 1e-3: + print(f"[rank {rank}] layer {layer_id}: expected {expected}, got {got}") + ok = False + assert ok, f"rank {rank} read stale/incorrect broadcast contents" + + # Indexer buffer owner-broadcast: owner writes a layer-distinct value, then + # every rank must read it back for every layer. + for layer_id in range(LAYER_NUM): + store_buf = pool.get_index_k_with_scale_buffer(layer_id) + assert ( + store_buf.data_ptr() + == pool.index_k_with_scale_buffer[layer_id].data_ptr() + ) + if pool._is_layer_owned(layer_id): + store_buf.fill_(layer_id + 10) + torch.cuda.synchronize() + torch.distributed.barrier() + for layer_id in range(LAYER_NUM): + # invalidate any cached remote copy so the read forces a fresh broadcast + pool.invalidate_index_buffer_for_layer(layer_id) + buf = pool._get_broadcastable_index_buffer(layer_id) + expected = layer_id + 10 + got = buf.float().mean().item() + if abs(got - expected) > 1e-3: + print( + f"[rank {rank}] index layer {layer_id}: exp {expected}, got {got}" + ) + ok = False + assert ok, f"rank {rank} read stale/incorrect index broadcast contents" + + # Async prefetch path: prefetch layer, then read must return owner value. + for layer_id in range(LAYER_NUM): + pool.remote_kv_layer_id = None # force a fresh broadcast + pool.prefetch_kv_buffer(layer_id) + buf = pool._get_broadcastable_kv_buffer(layer_id) + got = buf.float().mean().item() + if abs(got - float(layer_id + 1)) > 1e-3: + print(f"[rank {rank}] prefetch layer {layer_id}: got {got}") + ok = False + assert ok, f"rank {rank} prefetch path returned incorrect contents" + + print(f"[rank {rank}] OK: all {LAYER_NUM} layers read correct owner contents") + torch.distributed.barrier() class TestLayerSplitDSABroadcast(CustomTestCase): diff --git a/test/registered/unit/models/test_kimi_k25.py b/test/registered/unit/models/test_kimi_k25.py index 3e4a11f5694d..0c368c29bb1a 100644 --- a/test/registered/unit/models/test_kimi_k25.py +++ b/test/registered/unit/models/test_kimi_k25.py @@ -67,7 +67,6 @@ CudaIpcTensorTransportProxy, ) from sglang.srt.runtime_context import ( - ParallelContext, get_context, get_parallel, publish, @@ -911,7 +910,7 @@ def forward(self, pixel_values, _grid_thws): hot_items = pickle.loads(pickle.dumps(hot.mm_items)) with ( - patch.object(ParallelContext, "config", SimpleNamespace(tp_size=1)), + get_parallel().override(tp_size=1), patch( "sglang.srt.multimodal.processors.kimi_k25._gpu_preprocess_images", return_value=( @@ -975,7 +974,7 @@ def forward(self, pixel_values, _grid_thws): ) with ( - patch.object(ParallelContext, "config", SimpleNamespace(tp_size=1)), + get_parallel().override(tp_size=1), patch( "sglang.srt.multimodal.processors.kimi_k25._gpu_preprocess_images", return_value=(torch.full((1, 3), 2.0), torch.tensor([[1, 1, 1]])), diff --git a/test/registered/unit/models/test_zaya_cca.py b/test/registered/unit/models/test_zaya_cca.py index e9cfadd715f0..83fd2f1fdce2 100644 --- a/test/registered/unit/models/test_zaya_cca.py +++ b/test/registered/unit/models/test_zaya_cca.py @@ -29,6 +29,7 @@ import torch +from sglang.srt.runtime_context import get_context from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.layer_ut_utils import init_single_process_dist from sglang.test.test_utils import CustomTestCase @@ -36,10 +37,16 @@ register_cpu_ci(est_time=30, suite="base-a-test-cpu") -def _ensure_dist_initialized() -> None: +def _ensure_dist_initialized(cls) -> None: """CCA reads the TP rank / world size inside ``__init__`` to size its - head-parallel projections, so the groups must exist before construction.""" + head-parallel projections. The rank is the live group's, so the groups must + exist before construction; the size answers from the published ``parallel`` + bag, so the case has to publish a context as well. + """ init_single_process_dist() + override = get_context().override_server_args(tp_size=1) + override.install() + cls.addClassCleanup(override.restore) @dataclass(frozen=True) @@ -243,7 +250,7 @@ def _make_tiny_cca( class TestZayaCCA(CustomTestCase): @classmethod def setUpClass(cls) -> None: - _ensure_dist_initialized() + _ensure_dist_initialized(cls) def test_single_chunk_matches_reference(self): """A single-chunk extend with empty prefix matches the no-state path.""" @@ -540,7 +547,7 @@ class TestZayaCCATensorParallel(CustomTestCase): @classmethod def setUpClass(cls) -> None: - _ensure_dist_initialized() + _ensure_dist_initialized(cls) def _slice_full_state_dict_into_rank(self, ref_cca, tp_cca, tp_rank: int): """Copy the reference's full weights into the per-rank CCA, using the diff --git a/test/registered/unit/server_args/test_resolution_declarations.py b/test/registered/unit/server_args/test_resolution_declarations.py index 32e3e18ca8b3..6b2b33c2304e 100644 --- a/test/registered/unit/server_args/test_resolution_declarations.py +++ b/test/registered/unit/server_args/test_resolution_declarations.py @@ -433,19 +433,11 @@ def test_every_published_leaf_is_what_resolution_decided(self): mapping = namespace_of(ServerArgs) self.assertGreater(len(mapping), 400, "the namespace mapping collapsed") - # The five sizes keep a live property shadowing the bare name; the - # comparison below reaches them anyway, through `get_parallel().config`. - self.assertGreaterEqual( - _live_topology_leaves() - & { - "tp_size", - "pp_size", - "moe_dp_size", - "attn_cp_size", - "dcp_size", - }, - {"tp_size", "pp_size", "moe_dp_size", "attn_cp_size", "dcp_size"}, - "a parallel size stopped being served from the live topology", + self.assertEqual( + set(), + _live_topology_leaves() & set(mapping), + "a parallel leaf gained a live member of the same name, so the " + "comparison below reads the group rather than the published leaf", ) compared = 0 @@ -461,10 +453,6 @@ def test_every_published_leaf_is_what_resolution_decided(self): unreachable.append(f"no get_{groups[0]}() for {path}.{field}") continue node = accessor() - if groups[0] == "parallel": - # Bare names there are the live topology; the published - # leaves are one hop down, so the reader takes that hop. - node = node.config try: for group in groups[1:]: node = getattr(node, group) diff --git a/test/registered/unit/server_args/test_resolution_reads_the_declarations.py b/test/registered/unit/server_args/test_resolution_reads_the_declarations.py index 3d940ae483e9..ded92deff345 100644 --- a/test/registered/unit/server_args/test_resolution_reads_the_declarations.py +++ b/test/registered/unit/server_args/test_resolution_reads_the_declarations.py @@ -420,7 +420,7 @@ def reads_a_leaf_through_the_alias(runner): def hands_the_accessor_to_a_helper(): - return compute_world_size(get_server_args()) + return attention_backends_of(get_server_args()) def reads_the_view(runner): @@ -580,7 +580,7 @@ def test_the_scan_sees_every_spelling_that_reached_production(self): """ helpers = _config_reading_helpers() decided = _declared_fields() - for name in ("m3_fp8_attn_gemm_enabled", "compute_world_size"): + for name in ("m3_fp8_attn_gemm_enabled", "attention_backends_of"): self.assertIn(name, helpers, f"the helper derivation lost {name}") for field in ("speculative_num_draft_tokens", "attention_backend"): self.assertIn(field, decided, f"the declared set lost {field}") @@ -598,8 +598,8 @@ def test_the_scan_sees_every_spelling_that_reached_production(self): "sa_local.attention_backend", "self._server_args.attention_backend", "engine_args.attention_backend", - "compute_world_size(get_server_args())" - " reads " + ", ".join(helpers["compute_world_size"]), + "attention_backends_of(get_server_args())" + " reads " + ", ".join(helpers["attention_backends_of"]), }, "the scan lost a spelling, or started flagging a legal one:\n " + "\n ".join(sorted(flagged)), diff --git a/test/registered/unit/test_global_config_read_ratchet.py b/test/registered/unit/test_global_config_read_ratchet.py index b183374b63cf..fe4c1c07f94d 100644 --- a/test/registered/unit/test_global_config_read_ratchet.py +++ b/test/registered/unit/test_global_config_read_ratchet.py @@ -13,9 +13,6 @@ The reads that remain live in ``runtime_context.py`` (exempt by module): the ``@property`` / method members computed from several fields plus the HF config, which are not namespace leaves and have no home but ``ServerArgs``. -Separately, ``_CONFIGURED_SIZE_CALL_SITES`` registers every business read of -``get_parallel().config.`` — the config tier of a size whose bare name is -the live topology — with the reason the live property cannot serve it. What the scan sees: ``get_server_args().field``, an alias (``sa = get_server_args()`` then ``sa.field`` -- function-local, module-level, or parked @@ -48,250 +45,6 @@ # resolution pipeline. _SLOT_OWNERS = ("srt/runtime_context.py", "srt/server_args.py", "srt/arg_groups/") -# Every configured read of a live-shadowed size (``get_parallel().config.pp_size`` -# and its four siblings), with the reason the live topology cannot answer there. -# The test below asserts this map is exactly the set of such reads, so the -# reasons cannot drift away from the code. -_CONFIGURED_SIZE_CALL_SITES = { - ("srt/layers/cp/base.py", "attn_cp_size"): ( - "the lazy strategy bind in a worker: the CP group is what the strategy " - "is being built for, and the configured width is what describes it" - ), - ("benchmark/one_batch.py", "pp_size"): ( - "CPU affinity for this rank, computed right after the work function " - "publishes and before dist init, so the groups do not exist yet" - ), - ("benchmark/one_batch.py", "tp_size"): ( - "the same affinity computation: the layout is the configured one, and " - "the live group is not up at this point in the work function" - ), - ("srt/entrypoints/engine.py", "pp_size"): ( - "the launch path decides how many scheduler processes to spawn; it runs " - "before any of them exists, so there is no group to ask" - ), - ("srt/entrypoints/engine.py", "attn_cp_size"): ( - "the launcher's per-TP-rank layout, computed while deciding what to " - "spawn -- the groups it is laying out do not exist yet" - ), - ("srt/entrypoints/engine.py", "moe_dp_size"): ( - "the MoE factor of that same pre-spawn layout" - ), - ("srt/ray/engine.py", "pp_size"): ( - "the Ray driver sizes the actor placement group; the actors it is about " - "to create are the ones that will hold the process groups" - ), - ("srt/ray/engine.py", "tp_size"): ( - "the same placement arithmetic as the stage count: the driver sizes " - "the actors that will hold the process groups" - ), - ("srt/ray/data_parallel_controller.py", "tp_size"): ( - "the same arithmetic on the DP path, also in the driver" - ), - ("srt/ray/data_parallel_controller.py", "pp_size"): ( - "same placement arithmetic on the DP path -- ranks per TP group, " - "computed in the driver before the actors start" - ), - ("srt/ray/data_parallel_controller.py", "attn_cp_size"): ( - "the attention-CP factor of that same placement arithmetic, and the one " - "size whose live value cannot express the configured intent when " - "attn_cp_size > moe_dp_size aliases the groups" - ), - ("srt/layers/attention/dsa/dsa_indexer.py", "pp_size"): ( - "gates `pp_size > 1 and not get_pp_group()...`; the short circuit is the " - "point, since with PP off the group is never touched, which is what lets " - "the Indexer be constructed before distributed init" - ), - ("srt/managers/scheduler.py", "pp_size"): ( - "dispatch_event_loop picks the PP event loop; the MLX runner stub never " - "initializes torch.distributed, so the live property asserts before the " - "MLX loop can start -- the configured leaf answers the same value " - "wherever the live groups exist" - ), - ("srt/mem_cache/kv_cache_configurator.py", "pp_size"): ( - "decides whether the token capacity needs a cross-PP all-reduce at all; " - "asking the configured size keeps that decision independent of whether a " - "PP group is installed in this process" - ), - ("srt/layers/dp_attention.py", "attn_cp_size"): ( - "compared against the configured moe_dp_size below" - ), - ("srt/layers/dp_attention.py", "moe_dp_size"): ( - "the configuration this predicate detects (attn_cp_size > moe_dp_size) is " - "the one where initialize_model_parallel aliases _MOE_DP to _ATTN_CP, so " - "the live sizes are equal there and a live comparison is always false" - ), - ("srt/managers/scheduler.py", "tp_size"): ( - "configure_scheduler_process runs before the scheduler's own process " - "groups exist -- configuring the process is what it is for -- so there " - "is nothing live to ask yet" - ), - ("srt/managers/scheduler.py", "moe_dp_size"): ( - "same pre-distributed-init arithmetic in configure_scheduler_process" - ), - ("srt/managers/scheduler.py", "attn_cp_size"): ( - "same pre-distributed-init arithmetic in configure_scheduler_process" - ), - ("srt/managers/scheduler.py", "dcp_size"): ( - "same pre-distributed-init arithmetic in configure_scheduler_process" - ), - ("srt/model_executor/runner/base_runner.py", "tp_size"): ( - "the same window as the stage count next to it: a draft runner shares " - "the target's groups, so the live property would answer for the wrong " - "runner" - ), - ("srt/model_executor/cpu_graph_runner.py", "tp_size"): ( - "the same window, on the CPU graph path" - ), - ("srt/entrypoints/v1_loads.py", "tp_size"): ( - "the accelerator count is arithmetic over the launch shape, reported " - "from the tokenizer process, which holds no model groups" - ), - ("srt/disaggregation/nixl/conn.py", "tp_size"): ( - "the NIXL rank arithmetic runs on the transfer path, which the CPU-only " - "conn tests exercise without starting torch.distributed" - ), - ("srt/managers/tokenizer_control_mixin.py", "tp_size"): ( - "the tokenizer divides its worker count by the launch width; it holds " - "no model groups" - ), - ("srt/model_executor/runner/base_runner.py", "pp_size"): ( - "the runner's layer window is arithmetic over the configured stage " - "count; a draft runner shares the target's groups, so the live " - "property would answer for the wrong runner" - ), - ("srt/model_executor/cpu_graph_runner.py", "pp_size"): ( - "the same window, on the CPU graph path" - ), - ( - "srt/managers/scheduler_components/metrics_reporter.py", - "pp_size", - ): ( - "the reporter labels its metrics with the stage count it was launched " - "with, which is configuration; the live group answers per process" - ), - ("srt/speculative/eagle_draft_cuda_graph_runner.py", "pp_size"): ( - "the draft runner's window over the target's stages: its own groups are " - "the target's, so the configured count is the one that describes it" - ), - ( - "srt/speculative/eagle_draft_extend_cuda_graph_runner.py", - "pp_size", - ): ("the same draft window, on the extend path"), - ( - "srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py", - "pp_size", - ): ("the same draft window, multi-layer extend"), - ("srt/speculative/frozen_kv_mtp_cuda_graph_runner.py", "pp_size"): ( - "the same draft window, frozen-KV MTP" - ), - ("srt/managers/data_parallel_controller.py", "pp_size"): ( - "the controller lays out its schedulers' ranks before spawning them, so " - "the groups it is sizing for do not exist yet" - ), - ("srt/managers/data_parallel_controller.py", "attn_cp_size"): ( - "the same pre-spawn rank arithmetic" - ), - ("srt/managers/data_parallel_controller.py", "moe_dp_size"): ( - "the same pre-spawn rank arithmetic" - ), - ("srt/entrypoints/v1_loads.py", "pp_size"): ( - "the /v1/loads accelerator count is arithmetic over the launch shape, " - "reported from the tokenizer process, which holds no model groups" - ), - ("srt/disaggregation/common/conn.py", "pp_size"): ( - "the bootstrap connection is built by the KV manager on the transfer " - "path, which the CPU-only conn tests exercise without ever starting " - "torch.distributed" - ), - ("srt/elastic_ep/elastic_ep.py", "tp_size"): ( - "the joiner's rank window is computed against the size the process was " - "configured with, not the size of the group it is about to join" - ), - ("srt/elastic_ep/expert_backup_manager.py", "tp_size"): ( - "the backup server counts the clients it expects to report in, which " - "is how many the launch configured -- the live group is what they are " - "still joining" - ), - ( - "srt/model_executor/model_runner_components/startup_weight_load.py", - "tp_size", - ): ( - "the load options are assembled in ModelRunner.__init__ for a runner " - "that may be a draft, whose groups are the target's; the configured " - "sizes are what the record answered before" - ), - ( - "srt/model_executor/model_runner_components/startup_weight_load.py", - "pp_size", - ): ("same options object, same reason"), - ( - "srt/model_executor/model_runner_components/startup_weight_load.py", - "attn_cp_size", - ): ("same options object, same reason"), - ( - "srt/model_executor/model_runner_components/startup_weight_load.py", - "dcp_size", - ): ("same options object, same reason"), - ( - "srt/model_executor/model_runner_components/spec_aux_hidden_state.py", - "tp_size", - ): ( - "the draft KV bytes/token estimate sizes the memory pool before the " - "draft runner exists, so its shard count is configuration" - ), - ("srt/eplb/expert_location.py", "tp_size"): ( - "the elastic-EP joiner window, used to size the expert layout: the " - "size the process was configured with, not the group it is joining" - ), - ("srt/utils/cuda_vmm_transport_utils.py", "tp_size"): ( - "the consumer count is configured fan-out arithmetic (tp_size // " - "dp_size), which is what the record answered before" - ), - ("srt/disaggregation/encoder/runtime.py", "tp_size"): ( - "the encode server's launch entry sizes its workers before it has " - "spawned any of them" - ), - ("srt/disaggregation/encoder/grpc_server.py", "tp_size"): ( - "the same worker-count arithmetic on the gRPC entry: it spawns the TP " - "workers, so their groups do not exist yet" - ), - ("srt/disaggregation/encoder/server.py", "tp_size"): ( - "`MMEncoder` builds its own TP group from this size -- " - "`initialize_model_parallel` is the call being handed it, so there is " - "nothing live to ask" - ), - ("srt/disaggregation/encoder/receiver.py", "tp_size"): ( - "the receiver labels and shards by the launch width; it runs in the " - "tokenizer process, which holds no encoder groups" - ), - ("srt/managers/rust_server.py", "tp_size"): ( - "the rust server decides its transport from the launch width, in the " - "tokenizer process, which holds no model groups" - ), - ("compile_deep_gemm.py", "tp_size"): ( - "the warm-up request fans bootstrap rooms across the launch's ranks; it " - "runs in the tokenizer process, which holds no model groups" - ), - ("srt/utils/common.py", "tp_size"): ( - "the require_*_tp_gather predicates compared the configured tp_size " - "when they read the record; the live property answers a different " - "question wherever the groups alias, so the configured accessor is the " - "mechanical substitution and the live one would be a semantic change" - ), - ("srt/model_loader/loader.py", "moe_dp_size"): ( - "the same dict already carries the live moe_dp_size under 'dp'; this entry " - "is the configured intent" - ), - ("srt/models/kimi_k25.py", "tp_size"): ( - "the IPC refcount must match the configured TP consumer count captured " - "when the tokenizer creates MmItemMemoryPool; a live attention subgroup " - "size could strand leases in the bounded pool" - ), - ("srt/models/kimi_k3.py", "tp_size"): ( - "same as kimi_k25: the IPC refcount must agree with the recycler's waiter" - ), -} - _DIRECT_BASELINE = 0 _ALIAS_BASELINE = 0 @@ -594,222 +347,6 @@ def test_global_field_reads_match_the_baseline(self): self._check("alias-form", alias, _ALIAS_BASELINE) -def _live_shadowed_sizes() -> frozenset: - """Names that are BOTH a live ``ParallelContext`` property and a ``parallel`` - config leaf. - - Derived from the two sides themselves: a size that gains a live property, or - a live property that gains a leaf, joins the registry's subject set without a - list here. - """ - from sglang.srt.arg_groups.arg_utils import namespace_of - from sglang.srt.runtime_context import ParallelContext - from sglang.srt.server_args import ServerArgs - - live = { - name - for name, value in vars(ParallelContext).items() - if isinstance(value, property) - } - leaves = { - field for field, path in namespace_of(ServerArgs).items() if path == "parallel" - } - shadowed = frozenset(live & leaves) - assert shadowed, "no live-shadowed size found; the derivation is broken" - return shadowed - - -def _parallel_config_reads(tree, subjects): - """Names in ``subjects`` read through the parallel bag's ``config`` hop. - - Sees ``get_parallel().config.pp_size``, the module-qualified spelling, a - local bound to either hop (``p = get_parallel()`` / ``cfg = p.config``), and - the ``getattr`` form of each. - """ - fns, modules = set(), set() - for node in ast.walk(tree): - if ( - isinstance(node, ast.ImportFrom) - and node.module - and node.module.endswith("runtime_context") - ): - fns |= {a.asname or a.name for a in node.names if a.name == "get_parallel"} - elif isinstance(node, ast.ImportFrom) and node.module: - # `from sglang.srt import runtime_context as rc` binds the module. - for a in node.names: - if f"{node.module}.{a.name}".endswith("runtime_context"): - modules.add(a.asname or a.name) - elif isinstance(node, ast.Import): - for a in node.names: - if a.name.endswith("runtime_context"): - # Unaliased, the call site spells the whole dotted path. - modules.add(a.asname or a.name) - - def dotted(node): - parts = [] - while isinstance(node, ast.Attribute): - parts.append(node.attr) - node = node.value - if not isinstance(node, ast.Name): - return None - parts.append(node.id) - return ".".join(reversed(parts)) - - def is_bag_call(node): - if not isinstance(node, ast.Call): - return False - func = node.func - if isinstance(func, ast.Name): - return func.id in fns - return ( - isinstance(func, ast.Attribute) - and func.attr == "get_parallel" - and dotted(func.value) in modules - ) - - bag_aliases, config_aliases = set(), set() - for _ in range(2): # a local copy of a local is still the same object - for node in ast.walk(tree): - if not isinstance(node, ast.Assign): - continue - value = node.value - if is_bag_call(value) or ( - isinstance(value, ast.Name) and value.id in bag_aliases - ): - bucket = bag_aliases - elif ( - isinstance(value, ast.Attribute) - and value.attr == "config" - and ( - is_bag_call(value.value) - or ( - isinstance(value.value, ast.Name) - and value.value.id in bag_aliases - ) - ) - ) or (isinstance(value, ast.Name) and value.id in config_aliases): - bucket = config_aliases - else: - continue - bucket |= {t.id for t in node.targets if isinstance(t, ast.Name)} - - def is_config_hop(node): - return ( - isinstance(node, ast.Attribute) - and node.attr == "config" - and ( - is_bag_call(node.value) - or (isinstance(node.value, ast.Name) and node.value.id in bag_aliases) - ) - ) or (isinstance(node, ast.Name) and node.id in config_aliases) - - found = set() - for node in ast.walk(tree): - if isinstance(node, ast.Attribute) and node.attr in subjects: - base, name = node.value, node.attr - elif ( - isinstance(node, ast.Call) - and isinstance(node.func, ast.Name) - and node.func.id == "getattr" - and len(node.args) >= 2 - and isinstance(node.args[1], ast.Constant) - and node.args[1].value in subjects - ): - base, name = node.args[0], node.args[1].value - else: - continue - if is_config_hop(base): - found.add(name) - return found - - -_READ_SPELLINGS = ( - "from sglang.srt.runtime_context import get_parallel\nx = get_parallel().config.tp_size", - "from sglang.srt.runtime_context import get_parallel as gp\nx = gp().config.tp_size", - "from sglang.srt import runtime_context as rc\nx = rc.get_parallel().config.tp_size", - "import sglang.srt.runtime_context\nx = sglang.srt.runtime_context.get_parallel().config.tp_size", - "from sglang.srt.runtime_context import get_parallel\np = get_parallel()\nx = p.config.tp_size", - "from sglang.srt.runtime_context import get_parallel\nc = get_parallel().config\nx = c.tp_size", - 'from sglang.srt.runtime_context import get_parallel\nx = getattr(get_parallel().config, "tp_size")', -) - - -class TestParallelConfigReadSpellings(CustomTestCase): - """``_parallel_config_reads`` resolves every spelling it claims to. - - The scan below decides what the documented set is compared against, so a - spelling it cannot resolve does not fail anything -- it drops the read. - """ - - def test_every_documented_spelling_resolves(self): - for source in _READ_SPELLINGS: - with self.subTest(source=source): - found = _parallel_config_reads(ast.parse(source), {"tp_size"}) - self.assertEqual({"tp_size"}, set(found)) - - def test_the_live_property_is_not_a_config_read(self): - source = ( - "from sglang.srt.runtime_context import get_parallel\n" - "x = get_parallel().tp_size" - ) - self.assertEqual( - set(), set(_parallel_config_reads(ast.parse(source), {"tp_size"})) - ) - - -class TestConfiguredSizeCallSites(CustomTestCase): - """The configured-vs-live exceptions are enumerated, with reasons. - - ``get_parallel().config.tp_size`` answers what the process was configured - with where the bare ``get_parallel().tp_size`` answers what the process ended - up with. Each site that needs the former is listed above with why the live - property cannot serve it, and this case fails if the code and that list - disagree. - - The unit is **(file, size)**, not the individual read: a second - ``.config.pp_size`` in a file already registered for it collapses into the - same entry, so the reason has to cover the file's use of that size rather - than one line. A new file, or a new size in a listed file, is what this - catches -- through any spelling of the hop. - """ - - def test_the_call_sites_match_the_documented_set(self): - subjects = _live_shadowed_sizes() - found = set() - scanned = 0 - for path in sorted(_PACKAGE_ROOT.rglob("*.py")): - rel = path.relative_to(_PACKAGE_ROOT).as_posix() - if rel.startswith(_SLOT_OWNERS): - continue - source = path.read_text() - # Every spelling `_parallel_config_reads` resolves -- the direct - # call, an aliased import, a module-qualified call, a local bound to - # either hop -- needs the name in the source, so skipping the rest is - # free. Filtering on anything narrower silently empties the scan. - if "get_parallel" not in source: - continue - scanned += 1 - try: - tree = ast.parse(source) - except SyntaxError: - continue - found |= {(rel, name) for name in _parallel_config_reads(tree, subjects)} - self.assertGreater( - scanned, - 50, - f"the pre-filter left only {scanned} files to scan; the derivation " - "is broken, not the tree", - ) - documented = set(_CONFIGURED_SIZE_CALL_SITES) - self.assertEqual( - documented, - found, - "configured-size reads drifted from their documented reasons.\n" - f" undocumented: {sorted(found - documented)}\n" - f" stale entries: {sorted(documented - found)}", - ) - - class TestNoRenamedAccessorImports(CustomTestCase): """The baseline scanner matches ``get_server_args`` by its literal name, so an ``import ... as`` rename would walk a read straight past the zero diff --git a/test/registered/unit/test_launch_path_reads_configured_sizes.py b/test/registered/unit/test_launch_path_reads_configured_sizes.py deleted file mode 100644 index 359dcdc19d6e..000000000000 --- a/test/registered/unit/test_launch_path_reads_configured_sizes.py +++ /dev/null @@ -1,419 +0,0 @@ -"""Launch paths read the configured parallel sizes, not the live ones. - -`get_parallel().pp_size` and its four siblings are read-through properties over -the process groups, so they answer only after distributed init. The launcher -decides how many processes to spawn *before* that, and a live read there raises -`Distributed environment is not initialized` -- a startup crash no unit test -reaches, because nothing short of booting a server runs the launcher. The -configured answer is one hop away on the same object, -`get_parallel().config.pp_size`, which reads the published `parallel` bag. -""" - -import ast -import pathlib -import unittest - -import sglang -from sglang.test.ci.ci_register import register_cpu_ci -from sglang.test.test_utils import CustomTestCase - -register_cpu_ci(est_time=9, suite="base-a-test-cpu") - -_PACKAGE_ROOT = pathlib.Path(sglang.__file__).resolve().parent - - -def _live_shadowed() -> dict: - """{name: remedy} for every name that is BOTH a live ParallelContext - property and a `parallel` config leaf. - - Derived from the two sides themselves, so a new size that gains a live - property (or a live property that gains a leaf) is watched without a second - list here. ParallelContext shadows more properties than these -- every - `_v(name, ...)` one raises the same "Distributed environment is not - initialized" -- but only a shadowed name has a configured answer to point a - launcher at. - """ - from sglang.srt.arg_groups.arg_utils import namespace_of - from sglang.srt.runtime_context import ParallelContext - from sglang.srt.server_args import ServerArgs - - live = { - name - for name, value in vars(ParallelContext).items() - if isinstance(value, property) - } - leaves = { - field for field, path in namespace_of(ServerArgs).items() if path == "parallel" - } - shadowed = live & leaves - assert shadowed, ( - "no live-shadowed parallel size found; the derivation is broken, not " - "the tree" - ) - return {name: f"get_parallel().config.{name}" for name in sorted(shadowed)} - - -_LIVE_SHADOWED = _live_shadowed() - -# Launch paths that decide how many children to spawn are derived below -# from the spawn itself. These launch without a size-driven spawn, so no -# derivation reaches them and they are carried by hand. -_HAND_CARRIED = ( - "srt/entrypoints/http_server.py", - "srt/entrypoints/sidecar.py", - "srt/ray/data_parallel_controller.py", - "srt/ray/engine.py", - "srt/ray/http_server.py", -) - - -def _multiprocessing_names(tree): - """Names bound to multiprocessing, to one of its start contexts, or to the - process constructors themselves.""" - modules, constructors = set(), set() - for node in ast.walk(tree): - if isinstance(node, ast.Import): - for a in node.names: - if a.name == "multiprocessing" or a.name.startswith("multiprocessing."): - modules.add(a.asname or a.name.split(".")[0]) - elif a.name == "torch.multiprocessing": - modules.add(a.asname or "torch") - elif isinstance(node, ast.ImportFrom): - if node.module in ( - "multiprocessing", - "multiprocessing.context", - "torch.multiprocessing", - ): - constructors |= { - a.asname or a.name for a in node.names if a.name == "Process" - } - elif node.module == "concurrent.futures": - constructors |= { - a.asname or a.name - for a in node.names - if a.name == "ProcessPoolExecutor" - } - for node in ast.walk(tree): - if isinstance(node, ast.Assign) and isinstance(node.value, ast.Call): - func = node.value.func - if ( - isinstance(func, ast.Attribute) - and func.attr == "get_context" - and isinstance(func.value, ast.Name) - and func.value.id in modules - ): - modules |= {t.id for t in node.targets if isinstance(t, ast.Name)} - return modules, constructors - - -def _spawns_from_a_size(tree) -> bool: - """Does a function here spawn a child *and* read a live-shadowed size? - - Both tiers count: deriving on the live read alone drops a launcher from the - scan the moment it is converted, so the guard would only watch the ones - that already fail it. - """ - modules, constructors = _multiprocessing_names(tree) - for fn in ast.walk(tree): - if not isinstance(fn, (ast.FunctionDef, ast.AsyncFunctionDef)): - continue - spawns = False - for node in ast.walk(fn): - if not isinstance(node, ast.Call): - continue - func = node.func - if isinstance(func, ast.Attribute) and func.attr in ( - "Process", - "ProcessPoolExecutor", - "Popen", - "spawn", - ): - # `mp.Process(`, `mp.get_context("spawn").Process(` and - # `subprocess.Popen(` all reach a child process; the receiver of - # a chained call is itself a call, so this cannot require a bare - # Name. - spawns = True - elif isinstance(func, ast.Name) and func.id in constructors: - spawns = True - if not spawns: - continue - # A record read (`server_args.tp_size`) sizes a spawn too, but it cannot - # raise pre-dist; only a bag read is this guard's subject. - live, configured = _shadowed_size_reads(tree, scope=fn) - if live or configured: - return True - return False - - -def _parallel_bag_names(tree): - """What this module calls `get_parallel`, plus any runtime_context alias. - - A literal-name match reads only one spelling; an aliased import or a - module-qualified call is the same read with a different surface. - """ - names, modules = set(), set() - for node in ast.walk(tree): - if ( - isinstance(node, ast.ImportFrom) - and node.module - and node.module.endswith("runtime_context") - ): - names |= { - a.asname or a.name for a in node.names if a.name == "get_parallel" - } - elif isinstance(node, ast.ImportFrom) and node.module: - # `from sglang.srt import runtime_context as rc` binds the module, - # so `rc.get_parallel()` is the same call under another spelling. - for a in node.names: - if f"{node.module}.{a.name}".endswith("runtime_context"): - modules.add(a.asname or a.name) - elif isinstance(node, ast.Import): - for a in node.names: - if a.name.endswith("runtime_context"): - modules.add(a.asname or a.name.split(".")[0]) - return names, modules - - -def _is_parallel_bag_call(node, names, modules) -> bool: - if not isinstance(node, ast.Call): - return False - if isinstance(node.func, ast.Name): - return node.func.id in names - return ( - isinstance(node.func, ast.Attribute) - and node.func.attr == "get_parallel" - and isinstance(node.func.value, ast.Name) - and node.func.value.id in modules - ) - - -def _bag_aliases(tree, names, qualified): - """Locals bound to either tier: `p = get_parallel()` then `p.pp_size` is the - same live read one line later, and `cfg = get_parallel().config` then - `cfg.pp_size` is the same configured read.""" - live, config = set(), set() - for node in ast.walk(tree): - if not isinstance(node, ast.Assign): - continue - value = node.value - if _is_parallel_bag_call(value, names, qualified): - bucket = live - elif ( - isinstance(value, ast.Attribute) - and value.attr == "config" - and _is_parallel_bag_call(value.value, names, qualified) - ): - bucket = config - else: - continue - bucket |= {t.id for t in node.targets if isinstance(t, ast.Name)} - return live, config - - -def _shadowed_size_reads(module_tree, scope=None): - """(live, configured) reads of a live-shadowed size in `scope`. - - `.tp_size` is the live group; `.config.tp_size` - is the published leaf. Both spellings are reported so a caller can tell a - launcher that reads the topology at all from one that reads it live. - - What binds the bag -- the import, a module-level alias -- lives at module - scope, so those names always come from `module_tree` even when only one - function is being walked. Deriving them from the function alone finds no - import, reports no reads, and quietly answers "this launcher reads nothing". - """ - names, qualified = _parallel_bag_names(module_tree) - live_aliases, config_aliases = _bag_aliases(module_tree, names, qualified) - - def is_live_bag(node): - return _is_parallel_bag_call(node, names, qualified) or ( - isinstance(node, ast.Name) and node.id in live_aliases - ) - - def is_config_bag(node): - return ( - isinstance(node, ast.Attribute) - and node.attr == "config" - and is_live_bag(node.value) - ) or (isinstance(node, ast.Name) and node.id in config_aliases) - - live, configured = [], [] - for node in ast.walk(scope if scope is not None else module_tree): - if isinstance(node, ast.Attribute) and node.attr in _LIVE_SHADOWED: - base, name, spelling = node.value, node.attr, "attribute" - elif ( - isinstance(node, ast.Call) - and isinstance(node.func, ast.Name) - and node.func.id == "getattr" - and len(node.args) >= 2 - and isinstance(node.args[1], ast.Constant) - and node.args[1].value in _LIVE_SHADOWED - ): - base, name, spelling = node.args[0], node.args[1].value, "getattr" - else: - continue - if is_config_bag(base): - configured.append((node.lineno, name, spelling)) - elif is_live_bag(base): - live.append((node.lineno, name, spelling)) - return live, configured - - -def _launch_paths(): - """(relative path, tree) per module that runs before its process groups. - - A module that sizes a spawn loop from a parallel-bag size is derived from - the spawn itself; `_HAND_CARRIED` holds the launch entries that spawn - nothing, which no derivation can reach. - """ - seen = {} - sizes = frozenset(_LIVE_SHADOWED) - for path in sorted(_PACKAGE_ROOT.rglob("*.py")): - source = path.read_text() - # Every spawn shape below names Process, ProcessPoolExecutor or Popen. - if not any(name in source for name in ("Process", "Popen", "spawn")): - continue - if not any(name in source for name in sizes): - continue - try: - tree = ast.parse(source) - except SyntaxError: - continue - if _spawns_from_a_size(tree): - seen[str(path.relative_to(_PACKAGE_ROOT))] = tree - for rel in _HAND_CARRIED: - seen.setdefault(rel, ast.parse((_PACKAGE_ROOT / rel).read_text())) - return sorted(seen.items()) - - -class TestLaunchPathsReadConfiguredSizes(CustomTestCase): - def test_configured_sizes_hold_when_the_live_topology_disagrees(self): - """The other direction: groups exist and answer something else. - - The check above proves nobody reads a live size too early. It says - nothing about what `.config.` returns once the groups *are* up and - answering a different number -- which is not hypothetical: elastic EP - scales the live topology away from what the operator configured, and - that divergence is the entire reason the two tiers are separate. With - only the early-read direction covered, a `config` hop that quietly - delegated to the live property would look correct. - """ - import json - import os - import tempfile - from unittest.mock import patch - - from sglang.srt.runtime_context import ( - ParallelContext, - get_parallel, - publish, - reset_context, - ) - from sglang.srt.server_args import ServerArgs - - directory = tempfile.mkdtemp(prefix="configured_sizes_") - with open(os.path.join(directory, "config.json"), "w") as handle: - json.dump( - { - "architectures": ["LlamaForCausalLM"], - "model_type": "llama", - "hidden_size": 16, - "intermediate_size": 32, - "num_attention_heads": 2, - "num_key_value_heads": 2, - "num_hidden_layers": 2, - "vocab_size": 128, - "max_position_embeddings": 2048, - }, - handle, - ) - # No resolve_once() here: `tp_size` is raw input, so the configured - # value is 2 either way. - server_args = ServerArgs(model_path=directory, device="cuda", tp_size=2) - self.addCleanup(reset_context) - publish(server_args, role="scheduler") - - # The live getter behind each property, read out of ParallelContext - # rather than listed here. - context_source = ast.parse( - (_PACKAGE_ROOT / "srt" / "runtime_context.py").read_text( - encoding="utf-8-sig" - ) - ) - parallel_class = next( - node - for node in ast.walk(context_source) - if isinstance(node, ast.ClassDef) and node.name == "ParallelContext" - ) - live_getter = {} - for method in parallel_class.body: - if not isinstance(method, ast.FunctionDef): - continue - for call in ast.walk(method): - if not ( - isinstance(call, ast.Call) - and isinstance(call.func, ast.Attribute) - and call.func.attr == "_v" - and call.args - and isinstance(call.args[0], ast.Constant) - ): - continue - getter = call.args[1] - if isinstance(getter, ast.Attribute): - live_getter[call.args[0].value] = getter.attr - state = "sglang.srt.distributed.parallel_state" - missing = sorted(set(_LIVE_SHADOWED) - set(live_getter)) - self.assertEqual( - missing, - [], - f"these sizes no longer have a live property to diverge from: {missing}", - ) - for name in sorted(_LIVE_SHADOWED): - with self.subTest(size=name): - target = f"{state}.{live_getter[name]}" - configured = getattr(get_parallel().config, name) - with patch(target, return_value=configured + 41): - self.assertEqual( - get_parallel().__getattribute__(name), - configured + 41, - f"{name} no longer follows the live topology", - ) - self.assertEqual( - getattr(get_parallel().config, name), - configured, - f"get_parallel().config.{name} followed the live topology " - "instead of the published configuration", - ) - from sglang.srt.arg_groups.overrides import resolution_result - - self.assertEqual( - resolution_result(server_args, "nccl_port"), - getattr(get_parallel(), "nccl_port"), - "a config-only leaf read bare disagreed with what resolution decided", - ) - reset_context() - - with self.assertRaisesRegex(ValueError, r"'parallel' not published"): - getattr(ParallelContext(), "nccl_port") - with self.assertRaisesRegex(AttributeError, r"has no 'not_a_leaf'"): - getattr(ParallelContext(), "not_a_leaf") - - def test_no_live_topology_read_before_distributed_init(self): - offenders = [] - for rel, tree in _launch_paths(): - live, _ = _shadowed_size_reads(tree) - for lineno, name, spelling in live: - through = " through getattr" if spelling == "getattr" else "" - offenders.append( - f"{rel}:{lineno} reads the live {name}{through}; " - f"use {_LIVE_SHADOWED[name]}" - ) - self.assertEqual( - offenders, - [], - "launch paths run before distributed init:\n " + "\n ".join(offenders), - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/unit/test_ray_driver_reads_the_bags.py b/test/registered/unit/test_ray_driver_reads_the_bags.py index 7532334e7294..5b412148e926 100644 --- a/test/registered/unit/test_ray_driver_reads_the_bags.py +++ b/test/registered/unit/test_ray_driver_reads_the_bags.py @@ -61,7 +61,7 @@ def test_the_world_size_follows_a_post_publish_override(self): self._publish(tp_size=2, pp_size=1, dp_size=1, enable_dp_attention=False) self.assertEqual(_compute_world_size(), 2) get_context().override("test.ray_driver", tp_size=8) - self.assertEqual(get_parallel().config.tp_size, 8) + self.assertEqual(get_parallel().tp_size, 8) self.assertEqual(_compute_world_size(), 8) def test_the_driver_modules_read_no_field_off_a_record(self): @@ -112,8 +112,7 @@ def test_the_driver_modules_read_no_field_off_a_record(self): offenders, [], "the Ray driver reads a config field off a record; the driver runs " - "after the publish, so read `get_parallel().config`:\n " - + "\n ".join(offenders), + "after the publish, so read `get_parallel()`:\n " + "\n ".join(offenders), ) diff --git a/test/registered/unit/test_runtime_context.py b/test/registered/unit/test_runtime_context.py index fdddaef634d7..9e2c10136f77 100644 --- a/test/registered/unit/test_runtime_context.py +++ b/test/registered/unit/test_runtime_context.py @@ -39,21 +39,16 @@ SIZE_RANK_DELEGATIONS = [ ("world_size", f"{_PS}.get_world_size"), ("world_rank", f"{_PS}.get_world_rank"), - ("tp_size", f"{_PS}.get_tensor_model_parallel_world_size"), ("tp_rank", f"{_PS}.get_tensor_model_parallel_rank"), - ("dcp_size", f"{_PS}.get_dcp_world_size"), ("dcp_rank", f"{_PS}.get_dcp_rank"), - ("pp_size", f"{_PS}.get_pipeline_model_parallel_world_size"), ("pp_rank", f"{_PS}.get_pipeline_model_parallel_rank"), ("moe_ep_size", f"{_PS}.get_moe_expert_parallel_world_size"), ("moe_ep_rank", f"{_PS}.get_moe_expert_parallel_rank"), - ("moe_dp_size", f"{_PS}.get_moe_data_parallel_world_size"), ("moe_dp_rank", f"{_PS}.get_moe_data_parallel_rank"), ("moe_tp_size", f"{_PS}.get_moe_tensor_parallel_world_size"), ("moe_tp_rank", f"{_PS}.get_moe_tensor_parallel_rank"), ("attn_tp_size", f"{_PS}.get_attn_tensor_model_parallel_world_size"), ("attn_tp_rank", f"{_PS}.get_attn_tensor_model_parallel_rank"), - ("attn_cp_size", f"{_PS}.get_attn_context_model_parallel_world_size"), ("attn_cp_rank", f"{_PS}.get_attn_context_model_parallel_rank"), ("attn_dp_size", f"{_DP}.get_attention_dp_size"), ("attn_dp_rank", f"{_DP}.get_attention_dp_rank"), @@ -897,9 +892,8 @@ def probe(x): def test_parallel_config_leaves_trace_under_torch_compile(self): # Regression: gate helpers such as ``enable_moe_dense_fully_dp()`` read - # parallel config leaves inside compiled model forwards through the - # `config` property, which must stay dynamo-traceable - # (``object.__getattribute__`` graph-breaks). + # parallel config leaves inside compiled model forwards, which must + # stay dynamo-traceable (``object.__getattribute__`` graph-breaks). # fullgraph=True turns any graph break back into a failure. import torch @@ -1368,5 +1362,28 @@ def test_accessors_that_wrap_methods_call_them(self): self.assertEqual([], wrong, "\n".join(wrong)) +class TestParallelLeafReads(_IsolatedServerArgs): + """The contract ``ParallelContext.__getattr__`` answers a parallel leaf on.""" + + def test_a_leaf_answers_what_resolution_decided(self): + from sglang.srt.arg_groups.overrides import resolution_result + + with get_context().override_server_args() as server_args: + self.assertEqual( + resolution_result(server_args, "nccl_port"), + get_parallel().nccl_port, + "a parallel leaf read off the context disagreed with what " + "resolution decided", + ) + + def test_before_publish_the_error_names_the_namespace(self): + with self.assertRaisesRegex(ValueError, r"'parallel' not published"): + getattr(ParallelContext(), "nccl_port") + + def test_an_unknown_name_is_still_an_attribute_error(self): + with self.assertRaisesRegex(AttributeError, r"has no 'not_a_leaf'"): + getattr(ParallelContext(), "not_a_leaf") + + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/test_runtime_context_override.py b/test/registered/unit/test_runtime_context_override.py index fecea1ddb879..ce7504e57b1f 100644 --- a/test/registered/unit/test_runtime_context_override.py +++ b/test/registered/unit/test_runtime_context_override.py @@ -91,7 +91,7 @@ def test_set_internal_state_fields_reach_parallel_and_spec(self): speculative_accept_threshold_single=0.5, speculative_accept_threshold_acc=0.9, ) - self.assertEqual(rc.get_parallel().config.pp_max_micro_batch_size, 8) + self.assertEqual(rc.get_parallel().pp_max_micro_batch_size, 8) self.assertEqual(rc.get_spec().speculative_accept_threshold_single, 0.5) self.assertEqual(rc.get_spec().speculative_accept_threshold_acc, 0.9)