Merge /get_load into /v1/loads (#23010)

This commit is contained in:
Liangsheng Yin
2026-04-17 13:36:51 -07:00
committed by GitHub
parent 44e67c6835
commit 573e12a7fc
13 changed files with 235 additions and 84 deletions
+19 -2
View File
@@ -647,12 +647,29 @@ async def server_info():
@app.get("/get_load")
async def get_load():
"""Get load metrics (deprecated - use /v1/loads instead)."""
"""Get load metrics (deprecated - use /v1/loads instead).
Legacy shim backed by /v1/loads. Projects GetLoadsReqOutput down to the
historical field shape (dp_rank, num_reqs, num_waiting_reqs, num_tokens,
num_pending_tokens, ts_tic) so existing clients keep working.
"""
logger.warning(
"Endpoint '/get_load' is deprecated and will be removed in a future version. "
"Please use '/v1/loads' instead."
)
return await _global_state.tokenizer_manager.get_load()
load_results = await _global_state.tokenizer_manager.get_loads(include=["core"])
ts = time.perf_counter()
return [
{
"dp_rank": r.dp_rank,
"num_reqs": r.num_running_reqs + r.num_waiting_reqs,
"num_waiting_reqs": r.num_waiting_reqs,
"num_tokens": r.num_total_tokens,
"num_pending_tokens": r.num_total_tokens - r.num_used_tokens,
"ts_tic": ts,
}
for r in load_results
]
# example usage:
@@ -65,6 +65,8 @@ def _compute_aggregate(load_dicts: list) -> dict:
"total_running_reqs": 0,
"total_waiting_reqs": 0,
"total_reqs": 0,
"total_used_tokens": 0,
"total_tokens": 0,
"avg_token_usage": 0.0,
"avg_throughput": 0.0,
"avg_utilization": 0.0,
@@ -77,6 +79,8 @@ def _compute_aggregate(load_dicts: list) -> dict:
"total_reqs": sum(
d["num_running_reqs"] + d["num_waiting_reqs"] for d in load_dicts
),
"total_used_tokens": sum(d["num_used_tokens"] for d in load_dicts),
"total_tokens": sum(d["num_total_tokens"] for d in load_dicts),
"avg_token_usage": round(sum(d["token_usage"] for d in load_dicts) / n, 4),
"avg_throughput": round(sum(d["gen_throughput"] for d in load_dicts) / n, 2),
"avg_utilization": round(sum(d["utilization"] for d in load_dicts) / n, 4),
@@ -93,8 +93,10 @@ class DPBudget:
def update_budget(self, load_update: WatchLoadUpdateReq):
"""Update the budget."""
for load in load_update.loads:
self.total_requests[load.dp_rank] = load.num_reqs
self.total_tokens[load.dp_rank] = load.num_tokens
self.total_requests[load.dp_rank] = (
load.num_running_reqs + load.num_waiting_reqs
)
self.total_tokens[load.dp_rank] = load.num_total_tokens
def dispatch(self, method: LoadBalanceMethod):
if method == LoadBalanceMethod.TOTAL_REQUESTS:
+8 -18
View File
@@ -1087,7 +1087,7 @@ class BatchTokenIDOutput(BaseBatchReq, SpeculativeDecodingMetricsMixin):
token_steps: List[List[int]] = None
# Load for DP balance
load: GetLoadReqOutput = None
load: GetLoadsReqOutput = None
# Customized info
customized_info: Optional[Dict[str, List[Any]]] = None
# Detailed breakdown of cached tokens by source (device/host/storage)
@@ -1149,7 +1149,7 @@ class BatchStrOutput(BaseBatchReq, SpeculativeDecodingMetricsMixin):
token_steps: List[List[int]] = None
# Load for DP balance
load: GetLoadReqOutput = None
load: GetLoadsReqOutput = None
# Customized info
customized_info: Optional[Dict[str, List[Any]]] = None
@@ -1860,21 +1860,6 @@ class BlockReqInput(BaseReq):
type: BlockReqType
@dataclass
class GetLoadReqInput(BaseReq):
pass
@dataclass
class GetLoadReqOutput(BaseReq):
dp_rank: int
num_reqs: int
num_waiting_reqs: int
num_tokens: int
num_pending_tokens: int
ts_tic: float
@dataclass
class MemoryMetrics:
"""Memory breakdown metrics."""
@@ -1992,6 +1977,11 @@ class GetLoadsReqOutput(BaseReq):
num_used_tokens: int = field(
metadata={"metric": ("gauge", "Number of tokens in use")}
)
# num_used_tokens + pending prefill tokens (waiting-queue seqlen, incl.
# disagg bootstrap/prealloc/transfer queues). Used for DP balance.
num_total_tokens: int = field(
metadata={"metric": ("gauge", "Used tokens plus pending prefill tokens")}
)
max_total_num_tokens: int = field(
metadata={"metric": ("gauge", "Maximum token capacity")}
)
@@ -2020,7 +2010,7 @@ class GetLoadsReqOutput(BaseReq):
@dataclass
class WatchLoadUpdateReq(BaseReq):
loads: List[GetLoadReqOutput]
loads: List[GetLoadsReqOutput]
@dataclass
+1 -3
View File
@@ -109,7 +109,6 @@ from sglang.srt.managers.io_struct import (
FreezeGCReq,
GetInternalStateReq,
GetInternalStateReqOutput,
GetLoadReqInput,
GetLoadsReqInput,
GetWeightsByNameReqInput,
HealthCheckOutput,
@@ -1321,7 +1320,6 @@ class Scheduler(
self.load_lora_adapter_from_tensors,
),
(UnloadLoRAAdapterReqInput, self.unload_lora_adapter),
(GetLoadReqInput, self.get_load),
(GetLoadsReqInput, self.get_loads),
(PauseGenerationReqInput, self.pause_generation),
(ContinueGenerationReqInput, self.continue_generation),
@@ -2360,7 +2358,7 @@ class Scheduler(
# For prefill-only batch, filter out finished requests since they
# won't go through the decode step. This keeps running_batch accurate
# for load reporting (num_running_reqs via /get_load).
# for load reporting (num_running_reqs via /v1/loads).
# Runs outside the last_batch block so stale requests are cleaned
# even when no new batches arrive (e.g. traffic stops).
if self.running_batch.is_prefill_only:
@@ -13,6 +13,7 @@ from sglang.srt.managers.io_struct import (
AbortReq,
BatchEmbeddingOutput,
BatchTokenIDOutput,
GetLoadsReqInput,
)
from sglang.srt.managers.schedule_batch import (
BaseFinishReason,
@@ -964,7 +965,7 @@ class SchedulerOutputProcessorMixin:
spec_acceptance_histogram = []
retraction_counts = []
output_hidden_states = None
load = self.get_load()
load = self.get_loads(GetLoadsReqInput(include=["core"]))
routed_experts = None
customized_info = {}
@@ -39,8 +39,6 @@ from sglang.srt.managers.io_struct import (
FlushCacheReqOutput,
GetInternalStateReq,
GetInternalStateReqOutput,
GetLoadReqInput,
GetLoadReqOutput,
GetLoadsReqInput,
GetLoadsReqOutput,
GetWeightsByNameReqInput,
@@ -121,7 +119,6 @@ _COMMUNICATOR_SPECS = [
("set_internal_state", SetInternalStateReqOutput),
("expert_distribution", ExpertDistributionReqOutput),
("update_lora_adapter", LoRAUpdateOutput),
("get_load", GetLoadReqOutput, "watching"),
("get_loads", GetLoadsReqOutput, "watching"),
("dumper_control", DumperControlReqOutput),
]
@@ -804,11 +801,6 @@ class TokenizerControlMixin:
self.auto_create_handle_loop()
return await self.dumper_control_communicator(obj)
async def get_load(self: TokenizerManager) -> List[GetLoadReqOutput]:
self.auto_create_handle_loop()
req = GetLoadReqInput()
return await self.get_load_communicator(req)
async def get_loads(
self: TokenizerManager,
include: Optional[List[str]] = None,
@@ -12,8 +12,6 @@ from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import (
DisaggregationMetrics,
GetLoadReqInput,
GetLoadReqOutput,
GetLoadsReqInput,
GetLoadsReqOutput,
LoRAMetrics,
@@ -776,11 +774,11 @@ class SchedulerMetricsMixin:
Args:
chunk_deduct: extra tokens to subtract from the chunked request's
remaining count. At batch-scheduling time the current chunk
remaining count. At batch-scheduling time the current chunk
has been planned but ``prefix_indices`` does not yet include it,
so callers pass ``extend_input_len`` here. At query time
(``get_load``) ``prefix_indices`` is already up-to-date, so
the default 0 is correct.
so callers pass ``extend_input_len`` here. At load-reporting
time ``prefix_indices`` is already up-to-date, so the default
0 is correct.
"""
num_pending_tokens = sum(req.seqlen for req in self.waiting_queue)
if self.chunked_req is not None:
@@ -788,31 +786,6 @@ class SchedulerMetricsMixin:
num_pending_tokens += req.seqlen - len(req.prefix_indices) - chunk_deduct
return num_pending_tokens
def get_load(self: Scheduler, _: GetLoadReqInput = None) -> GetLoadReqOutput:
num_tokens, _ = self.get_pool_stats().get_kv_token_stats()
num_pending_tokens = self._get_num_pending_tokens()
# Tokens and request count in waiting queue, bootstrap queue, prealloc queue
waiting_queues = [self.waiting_queue]
if self.disaggregation_mode == DisaggregationMode.PREFILL:
waiting_queues.append(self.disagg_prefill_bootstrap_queue.queue)
elif self.disaggregation_mode == DisaggregationMode.DECODE:
waiting_queues.append(self.disagg_decode_prealloc_queue.queue)
waiting_queues.append(self.disagg_decode_transfer_queue.queue)
waiting_queues.append(self.disagg_decode_prealloc_queue.retracted_queue)
num_tokens += sum(req.seqlen for queue in waiting_queues for req in queue)
num_waiting_reqs = sum(len(queue) for queue in waiting_queues)
return GetLoadReqOutput(
dp_rank=self.dp_rank,
num_reqs=len(self.running_batch.reqs) + num_waiting_reqs,
num_waiting_reqs=num_waiting_reqs,
num_tokens=num_tokens,
num_pending_tokens=num_pending_tokens,
ts_tic=time.perf_counter(),
)
def get_loads(self: Scheduler, req: GetLoadsReqInput = None) -> GetLoadsReqOutput:
"""
Get comprehensive load metrics for /v1/loads endpoint.
@@ -841,6 +814,9 @@ class SchedulerMetricsMixin:
num_waiting_reqs = sum(len(queue) for queue in waiting_queues)
num_used_tokens, kv_token_usage = self.get_pool_stats().get_kv_token_stats()
num_total_tokens = num_used_tokens + sum(
req.seqlen for queue in waiting_queues for req in queue
)
memory = None
if include_all or "memory" in include:
@@ -925,6 +901,7 @@ class SchedulerMetricsMixin:
num_running_reqs=num_running_reqs,
num_waiting_reqs=num_waiting_reqs,
num_used_tokens=num_used_tokens,
num_total_tokens=num_total_tokens,
max_total_num_tokens=self.max_total_num_tokens,
token_usage=round(kv_token_usage, 4),
gen_throughput=round(self.stats.gen_throughput, 2),
+7 -7
View File
@@ -209,7 +209,7 @@ impl WorkerManager {
url: &str,
api_key: Option<&str>,
) -> isize {
let load_url = format!("{}/get_load", url);
let load_url = format!("{}/v1/loads?include=core", url);
let mut req = client.get(&load_url).timeout(REQUEST_TIMEOUT);
if let Some(key) = api_key {
req = req.bearer_auth(key);
@@ -217,12 +217,12 @@ impl WorkerManager {
match req.send().await {
Ok(r) if r.status().is_success() => match r.json::<Value>().await {
Ok(json) if json.is_array() => json
.as_array()
.unwrap()
.iter()
.filter_map(|e| e.get("num_tokens").and_then(|v| v.as_i64()))
.sum::<i64>() as isize,
Ok(json) => json
.get("aggregate")
.and_then(|a| a.get("total_tokens"))
.and_then(|v| v.as_i64())
.map(|n| n as isize)
.unwrap_or(-1),
_ => -1,
},
_ => -1,
@@ -30,7 +30,7 @@ class TestNpuApi(CustomTestCase):
"""Testcase: Verify that the basic functions of the API interfaces work properly and the returned parameters are consistent with the configurations.
[Test Category] Interface
[Test Target] /health; /health_generate; /ping; /model_info; /server_info; /get_load; /v1/models; /v1/models/{model:path}; /generate
[Test Target] /health; /health_generate; /ping; /model_info; /server_info; /v1/loads; /v1/models; /v1/models/{model:path}; /generate
"""
@classmethod
@@ -90,15 +90,18 @@ class TestNpuApi(CustomTestCase):
self.assertEqual(response.json()["model_path"], self.model)
self.assertEqual(response.json()["tokenizer_path"], self.model)
def test_api_get_load(self):
response = requests.get(f"{self.base_url}/get_load")
def test_api_v1_loads(self):
response = requests.get(f"{self.base_url}/v1/loads")
self.assertEqual(response.status_code, 200)
self.assertIsNone(response.json()[0]["rid"])
self.assertIsNone(response.json()[0]["http_worker_ipc"])
self.assertIsNone(response.json()[0]["dp_rank"])
self.assertGreaterEqual(response.json()[0]["num_reqs"], 0)
self.assertGreaterEqual(response.json()[0]["num_waiting_reqs"], 0)
self.assertGreaterEqual(response.json()[0]["num_tokens"], 0)
body = response.json()
self.assertIn("loads", body)
self.assertIn("aggregate", body)
self.assertGreaterEqual(len(body["loads"]), 1)
load = body["loads"][0]
self.assertGreaterEqual(load["num_running_reqs"], 0)
self.assertGreaterEqual(load["num_waiting_reqs"], 0)
self.assertGreaterEqual(load["num_used_tokens"], 0)
self.assertGreaterEqual(load["num_total_tokens"], 0)
def test_api_v1_models(self):
response = requests.get(f"{self.base_url}/v1/models")
@@ -0,0 +1,76 @@
"""Unit tests for /v1/loads _compute_aggregate.
Narrow scope: lock in the semantic of new aggregate keys added by this PR
(total_used_tokens vs total_tokens). Trivial helpers (dict filtering,
zero-init branch) are not covered they would just restate Python.
"""
import unittest
from sglang.srt.entrypoints.v1_loads import _compute_aggregate
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=1, suite="stage-a-test-cpu")
def _load(
*,
dp_rank=0,
running=0,
waiting=0,
used=0,
total=0,
token_usage=0.0,
throughput=0.0,
utilization=0.0,
):
return {
"dp_rank": dp_rank,
"num_running_reqs": running,
"num_waiting_reqs": waiting,
"num_used_tokens": used,
"num_total_tokens": total,
"token_usage": token_usage,
"gen_throughput": throughput,
"utilization": utilization,
}
class TestComputeAggregate(CustomTestCase):
def test_multi_dp_rank_sums(self):
agg = _compute_aggregate(
[
_load(dp_rank=0, running=3, waiting=1, used=50, total=70),
_load(dp_rank=1, running=5, waiting=2, used=80, total=100),
_load(dp_rank=2, running=0, waiting=4, used=0, total=40),
]
)
self.assertEqual(agg["total_running_reqs"], 8)
self.assertEqual(agg["total_waiting_reqs"], 7)
self.assertEqual(agg["total_reqs"], 15)
self.assertEqual(agg["total_used_tokens"], 130)
self.assertEqual(agg["total_tokens"], 210)
def test_averages_over_dp_count(self):
agg = _compute_aggregate(
[
_load(token_usage=0.6, throughput=100.0, utilization=0.5),
_load(token_usage=0.8, throughput=200.0, utilization=0.7),
]
)
self.assertAlmostEqual(agg["avg_token_usage"], 0.7)
self.assertAlmostEqual(agg["avg_throughput"], 150.0)
self.assertAlmostEqual(agg["avg_utilization"], 0.6)
def test_total_tokens_differs_from_total_used_tokens(self):
# Regression: total_tokens sums num_total_tokens, NOT num_used_tokens.
# Gateway reads aggregate.total_tokens for DP load estimation, so a
# silent swap would under-report load.
agg = _compute_aggregate([_load(used=10, total=30), _load(used=20, total=45)])
self.assertEqual(agg["total_used_tokens"], 30)
self.assertEqual(agg["total_tokens"], 75)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,91 @@
"""Unit tests for DPBudget — field mapping regression guard.
This PR changed DPBudget.update_budget to read num_running_reqs +
num_waiting_reqs and num_total_tokens from the new GetLoadsReqOutput.
These tests lock in that mapping. Pre-existing dispatch logic is not
retested here it's covered by DP balance integration tests.
"""
import dataclasses
import unittest
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel
maybe_stub_sgl_kernel()
from sglang.srt.managers.data_parallel_controller import DPBudget
from sglang.srt.managers.io_struct import GetLoadsReqOutput, WatchLoadUpdateReq
register_cpu_ci(est_time=2, suite="stage-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,
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)
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),
]
)
)
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),
]
)
)
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,
)
]
)
)
self.assertEqual(budget.total_requests, [10, 2, 30])
self.assertEqual(budget.total_tokens, [100, 50, 300])
if __name__ == "__main__":
unittest.main()
@@ -33,7 +33,7 @@ class TestTypeBasedDispatcher(unittest.TestCase):
FlushCacheReqInput,
FreezeGCReq,
GetInternalStateReq,
GetLoadReqInput,
GetLoadsReqInput,
GetWeightsByNameReqInput,
InitWeightsSendGroupForRemoteInstanceReqInput,
InitWeightsUpdateGroupReqInput,
@@ -113,7 +113,7 @@ class TestTypeBasedDispatcher(unittest.TestCase):
(ExpertDistributionReq, lambda req: "expert_distribution_handled"),
(LoadLoRAAdapterReqInput, lambda req: "load_lora_adapter_handled"),
(UnloadLoRAAdapterReqInput, lambda req: "unload_lora_adapter_handled"),
(GetLoadReqInput, lambda req: "get_load_handled"),
(GetLoadsReqInput, lambda req: "get_loads_handled"),
]
# Create requests that conforms to the real distribution
@@ -204,7 +204,7 @@ class TestTypeBasedDispatcher(unittest.TestCase):
test_requests.append(GetWeightsByNameReqInput(name=""))
test_requests.append(ReleaseMemoryOccupationReqInput())
test_requests.append(RpcReqInput(method=""))
test_requests.append(GetLoadReqInput())
test_requests.append(GetLoadsReqInput())
dispatcher = TypeBasedDispatcher(mapping)