Optimize get load calls (/v1/loads) using shared-memory load snapshots (#26348)

Co-authored-by: cctry <cctry@meta.com>
This commit is contained in:
Lianmin Zheng
2026-05-29 13:40:26 -07:00
committed by GitHub
co-authored by cctry
parent 3cecc77ccb
commit 4ff1296f5e
14 changed files with 1521 additions and 271 deletions
@@ -11,11 +11,12 @@ if a scheduler starts reading another attr. `maybe_external_dp_rank_routing`
is exercised as the real method, no mock.
"""
import dataclasses
import unittest
from types import SimpleNamespace
from unittest.mock import MagicMock
import msgspec.structs
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel
@@ -26,29 +27,20 @@ from sglang.srt.managers.data_parallel_controller import (
DPBudget,
LoadBalanceMethod,
)
from sglang.srt.managers.io_struct import GetLoadsReqOutput, WatchLoadUpdateReq
from sglang.srt.managers.load_snapshot import LoadSnapshot
register_cpu_ci(est_time=11, suite="base-a-test-cpu")
_BASE_LOAD = GetLoadsReqOutput(
dp_rank=0,
timestamp=0.0,
num_running_reqs=0,
num_waiting_reqs=0,
num_used_tokens=0,
num_total_tokens=0,
_BASE_LOAD = msgspec.structs.replace(
LoadSnapshot(dp_rank=0),
max_total_num_tokens=4096,
token_usage=0.0,
gen_throughput=0.0,
cache_hit_rate=0.0,
utilization=0.0,
max_running_requests=128,
)
def _load(**overrides) -> GetLoadsReqOutput:
return dataclasses.replace(_BASE_LOAD, **overrides)
def _load(**overrides) -> LoadSnapshot:
return msgspec.structs.replace(_BASE_LOAD, **overrides)
def _make_controller(dp_size: int) -> DataParallelController:
@@ -74,44 +66,52 @@ class TestDPBudgetUpdateBudget(CustomTestCase):
def test_maps_running_plus_waiting_to_total_requests(self):
budget = DPBudget(dp_size=2)
budget.update_budget(
WatchLoadUpdateReq(
loads=[
_load(dp_rank=0, num_running_reqs=3, num_waiting_reqs=2),
_load(dp_rank=1, num_running_reqs=5, num_waiting_reqs=1),
]
)
[
_load(dp_rank=0, timestamp=1.0, num_running_reqs=3, num_waiting_reqs=2),
_load(dp_rank=1, timestamp=1.0, num_running_reqs=5, num_waiting_reqs=1),
]
)
self.assertEqual(budget.total_requests, [5, 6])
def test_maps_num_total_tokens_not_num_used_tokens(self):
# Reads num_total_tokens (used + pending prefill), NOT num_used_tokens.
# A silent swap here would break DP balance for long-prompt workloads.
budget = DPBudget(dp_size=2)
budget.update_budget(
WatchLoadUpdateReq(
loads=[
_load(dp_rank=0, num_used_tokens=100, num_total_tokens=150),
_load(dp_rank=1, num_used_tokens=80, num_total_tokens=80),
]
)
[
_load(
dp_rank=0, timestamp=1.0, num_used_tokens=100, num_total_tokens=150
),
_load(
dp_rank=1, timestamp=1.0, num_used_tokens=80, num_total_tokens=80
),
]
)
self.assertEqual(budget.total_tokens, [150, 80])
def test_partial_update_only_affects_reported_rank(self):
budget = DPBudget(dp_size=3)
budget.total_requests = [10, 20, 30]
budget.total_tokens = [100, 200, 300]
budget.update_budget(
WatchLoadUpdateReq(
loads=[
_load(
dp_rank=1,
num_running_reqs=1,
num_waiting_reqs=1,
num_total_tokens=50,
)
]
)
[
_load(
dp_rank=0, timestamp=1.0, num_running_reqs=10, num_total_tokens=100
),
_load(
dp_rank=1, timestamp=1.0, num_running_reqs=20, num_total_tokens=200
),
_load(
dp_rank=2, timestamp=1.0, num_running_reqs=30, num_total_tokens=300
),
]
)
budget.update_budget(
[
_load(
dp_rank=1,
timestamp=2.0,
num_running_reqs=1,
num_waiting_reqs=1,
num_total_tokens=50,
)
]
)
self.assertEqual(budget.total_requests, [10, 2, 30])
self.assertEqual(budget.total_tokens, [100, 50, 300])