From 573e12a7fc8bebde2f6caf9ddc5c89e0459ef4df Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Fri, 17 Apr 2026 13:36:51 -0700 Subject: [PATCH] Merge /get_load into /v1/loads (#23010) --- python/sglang/srt/entrypoints/http_server.py | 21 ++++- python/sglang/srt/entrypoints/v1_loads.py | 4 + .../srt/managers/data_parallel_controller.py | 6 +- python/sglang/srt/managers/io_struct.py | 26 ++---- python/sglang/srt/managers/scheduler.py | 4 +- .../scheduler_output_processor_mixin.py | 3 +- .../srt/managers/tokenizer_control_mixin.py | 8 -- .../observability/scheduler_metrics_mixin.py | 39 ++------ sgl-model-gateway/src/core/worker_manager.rs | 14 +-- .../ascend/interface/test_npu_api.py | 21 +++-- .../entrypoints/test_v1_loads_aggregate.py | 76 ++++++++++++++++ .../unit/managers/test_dp_budget.py | 91 +++++++++++++++++++ .../utils/test_type_based_dispatcher.py | 6 +- 13 files changed, 235 insertions(+), 84 deletions(-) create mode 100644 test/registered/unit/entrypoints/test_v1_loads_aggregate.py create mode 100644 test/registered/unit/managers/test_dp_budget.py diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index aba143c83..e2f005337 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -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: diff --git a/python/sglang/srt/entrypoints/v1_loads.py b/python/sglang/srt/entrypoints/v1_loads.py index 784417010..0b5bb6f5c 100644 --- a/python/sglang/srt/entrypoints/v1_loads.py +++ b/python/sglang/srt/entrypoints/v1_loads.py @@ -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), diff --git a/python/sglang/srt/managers/data_parallel_controller.py b/python/sglang/srt/managers/data_parallel_controller.py index 0c7921e92..3ac82063f 100644 --- a/python/sglang/srt/managers/data_parallel_controller.py +++ b/python/sglang/srt/managers/data_parallel_controller.py @@ -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: diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index 94a26966e..a3f6e4222 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -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 diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index bf7d45879..997903bf6 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -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: diff --git a/python/sglang/srt/managers/scheduler_output_processor_mixin.py b/python/sglang/srt/managers/scheduler_output_processor_mixin.py index 087aad6cf..27cd17025 100644 --- a/python/sglang/srt/managers/scheduler_output_processor_mixin.py +++ b/python/sglang/srt/managers/scheduler_output_processor_mixin.py @@ -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 = {} diff --git a/python/sglang/srt/managers/tokenizer_control_mixin.py b/python/sglang/srt/managers/tokenizer_control_mixin.py index b18c3cd34..c99999f4b 100644 --- a/python/sglang/srt/managers/tokenizer_control_mixin.py +++ b/python/sglang/srt/managers/tokenizer_control_mixin.py @@ -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, diff --git a/python/sglang/srt/observability/scheduler_metrics_mixin.py b/python/sglang/srt/observability/scheduler_metrics_mixin.py index ffc058cdf..30a8a6802 100644 --- a/python/sglang/srt/observability/scheduler_metrics_mixin.py +++ b/python/sglang/srt/observability/scheduler_metrics_mixin.py @@ -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), diff --git a/sgl-model-gateway/src/core/worker_manager.rs b/sgl-model-gateway/src/core/worker_manager.rs index 4a21ae3fb..28ef763d8 100644 --- a/sgl-model-gateway/src/core/worker_manager.rs +++ b/sgl-model-gateway/src/core/worker_manager.rs @@ -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::().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::() 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, diff --git a/test/registered/ascend/interface/test_npu_api.py b/test/registered/ascend/interface/test_npu_api.py index 81eaef4f8..e598e8bea 100644 --- a/test/registered/ascend/interface/test_npu_api.py +++ b/test/registered/ascend/interface/test_npu_api.py @@ -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") diff --git a/test/registered/unit/entrypoints/test_v1_loads_aggregate.py b/test/registered/unit/entrypoints/test_v1_loads_aggregate.py new file mode 100644 index 000000000..75e1cf50c --- /dev/null +++ b/test/registered/unit/entrypoints/test_v1_loads_aggregate.py @@ -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() diff --git a/test/registered/unit/managers/test_dp_budget.py b/test/registered/unit/managers/test_dp_budget.py new file mode 100644 index 000000000..ccf2f8449 --- /dev/null +++ b/test/registered/unit/managers/test_dp_budget.py @@ -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() diff --git a/test/registered/utils/test_type_based_dispatcher.py b/test/registered/utils/test_type_based_dispatcher.py index 37480a75f..9763f5188 100644 --- a/test/registered/utils/test_type_based_dispatcher.py +++ b/test/registered/utils/test_type_based_dispatcher.py @@ -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)