From e586a6f2c5f2d1e0626bbe0cb1580d56c12398a2 Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Mon, 24 Aug 2026 20:21:45 +0800 Subject: [PATCH] Report the whole server's world size in the scheduler's internal state (#35929) --- python/sglang/srt/managers/scheduler.py | 4 +- python/sglang/srt/server_args.py | 9 ++ .../test_scheduler_internal_state_env_vars.py | 4 + ...est_scheduler_internal_state_world_size.py | 128 ++++++++++++++++++ 4 files changed, 144 insertions(+), 1 deletion(-) create mode 100644 test/registered/unit/managers/test_scheduler_internal_state_world_size.py diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 8baa6be6e..3c0502a1f 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -44,6 +44,7 @@ from sglang.srt.runtime_context import ( get_observability, get_parallel, get_schedule, + get_server_args, get_serving, get_spec, ) @@ -296,7 +297,7 @@ from sglang.srt.plugins import load_plugins from sglang.srt.runtime_context import get_context, publish from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.sampling.sampling_params import TOP_K_ALL -from sglang.srt.server_args import PortArgs, ServerArgs +from sglang.srt.server_args import PortArgs, ServerArgs, compute_world_size from sglang.srt.session.session_controller import SessionController from sglang.srt.speculative.base_spec_worker import BaseSpecWorker from sglang.srt.speculative.dflash_utils import validate_dflash_request @@ -4421,6 +4422,7 @@ class Scheduler( # 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_server_args()) 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 diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index f50cf13be..a1ebc0cfc 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -10450,6 +10450,15 @@ class ServerArgs: return self.expert_balancedness_report_mode in ("prometheus", "both") +def compute_world_size(server_args: ServerArgs) -> int: + """Return the total GPU count across all data-parallel replicas.""" + return ( + (1 if server_args.enable_dp_attention else server_args.dp_size) + * server_args.tp_size + * server_args.pp_size + ) + + def m3_fp8_attn_gemm_enabled(args) -> bool: """Whether MiniMax-M3 attention GEMMs run in fp8 (no opt-in flag; active whenever possible): fp8_e4m3 main + index KV caches, fp8-cast q, fp8 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 6fdd6eaad..48606f9ba 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 @@ -48,6 +48,10 @@ class TestSchedulerInternalStateEnvVars(unittest.TestCase): ), patch( "sglang.srt.managers.scheduler.get_exec", return_value=SimpleNamespace(moe=SimpleNamespace(elastic_ep_backend=None)), + ), patch( + "sglang.srt.managers.scheduler.get_server_args", return_value=None + ), patch( + "sglang.srt.managers.scheduler.compute_world_size", return_value=1 ): output = scheduler.get_internal_state(recv_req=GetInternalStateReq()) 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 new file mode 100644 index 000000000..e4c66f478 --- /dev/null +++ b/test/registered/unit/managers/test_scheduler_internal_state_world_size.py @@ -0,0 +1,128 @@ +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 + +maybe_stub_sgl_kernel() + +from sglang.srt.managers.io_struct import GetInternalStateReq +from sglang.srt.managers.scheduler import Scheduler +from sglang.srt.server_args import ServerArgs, compute_world_size + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + + +def _make_server_args( + *, tp_size: int, pp_size: int, dp_size: int, enable_dp_attention: bool +) -> ServerArgs: + return ServerArgs( + model_path="dummy", + 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.""" + server_args = _make_server_args( + tp_size=1, pp_size=1, dp_size=1, enable_dp_attention=False + ) + + self.assertEqual(compute_world_size(server_args), 1) + + def test_tensor_and_pipeline_stages_multiply(self): + """Each (pp_rank, tp_rank) pair is its own scheduler process on its own gpu.""" + server_args = _make_server_args( + tp_size=2, pp_size=3, dp_size=1, enable_dp_attention=False + ) + + self.assertEqual(compute_world_size(server_args), 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.""" + server_args = _make_server_args( + tp_size=2, pp_size=1, dp_size=2, enable_dp_attention=False + ) + + self.assertEqual(compute_world_size(server_args), 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.""" + server_args = _make_server_args( + tp_size=4, pp_size=1, dp_size=2, enable_dp_attention=True + ) + + self.assertEqual(compute_world_size(server_args), 4) + + +class TestSchedulerInternalStateWorldSize(unittest.TestCase): + def _get_internal_state(self, server_args: ServerArgs) -> dict: + scheduler = Scheduler.__new__(Scheduler) + scheduler.metrics_reporter = SimpleNamespace( + last_gen_throughput=1.0, + spec_total_num_forward_ct=0, + spec_total_num_accept_tokens=0, + step_time_dict={}, + ) + scheduler.tp_worker = SimpleNamespace( + model_runner=SimpleNamespace(weight_load_mem_usage=1.0), + graph_memory_usage=None, + ) + scheduler.token_to_kv_pool_allocator = SimpleNamespace( + get_kvcache=lambda: SimpleNamespace(mem_usage=3.0) + ) + scheduler.startup_available_gpu_memory_gb = 4.0 + scheduler.startup_time = 1.0 + scheduler.max_total_num_tokens = 100 + scheduler.swa_tokens_per_layer = None + scheduler.max_running_requests = 8 + scheduler.spec_algorithm = SimpleNamespace( + is_none=lambda: True, + is_dspark=lambda: False, + ) + 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_server_args", + return_value=server_args, + ): + 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.""" + server_args = _make_server_args( + tp_size=2, pp_size=1, dp_size=2, enable_dp_attention=False + ) + + internal_state = self._get_internal_state(server_args) + + 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.""" + server_args = _make_server_args( + tp_size=2, pp_size=1, dp_size=2, enable_dp_attention=False + ) + + internal_state = self._get_internal_state(server_args) + + self.assertNotEqual( + internal_state["world_size"], server_args.tp_size * server_args.pp_size + ) + + +if __name__ == "__main__": + unittest.main()