Merge /get_load into /v1/loads (#23010)
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user