Report the whole server's world size in the scheduler's internal state (#35929)

This commit is contained in:
fzyzcjy
2026-08-24 20:21:45 +08:00
committed by GitHub
parent 6dd79576cd
commit e586a6f2c5
4 changed files with 144 additions and 1 deletions
+3 -1
View File
@@ -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
+9
View File
@@ -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
@@ -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())
@@ -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()