Fix get_load API (#13991)

This commit is contained in:
Liangsheng Yin
2025-11-26 22:07:51 +08:00
committed by GitHub
parent b704b0a922
commit eff6a07c8f
5 changed files with 41 additions and 81 deletions

View File

@@ -691,7 +691,8 @@ class Req:
self.dllm_config = dllm_config
@property
def seqlen(self):
def seqlen(self) -> int:
"""Get the current sequence length of the request."""
return len(self.origin_input_ids) + len(self.output_ids)
@property

View File

@@ -85,7 +85,6 @@ from sglang.srt.managers.io_struct import (
GetInternalStateReq,
GetInternalStateReqOutput,
GetLoadReqInput,
GetLoadReqOutput,
GetWeightsByNameReqInput,
HealthCheckOutput,
InitWeightsSendGroupForRemoteInstanceReqInput,
@@ -427,7 +426,6 @@ class Scheduler(
# Init running status
self.waiting_queue: List[Req] = []
self.num_waiting_tokens = 0
# The running decoding batch for continuous batching
self.running_batch: ScheduleBatch = ScheduleBatch(reqs=[], batch_is_full=False)
# The current forward batch
@@ -1465,7 +1463,6 @@ class Scheduler(
return
self._prefetch_kvcache(req)
self.waiting_queue.append(req)
self.num_waiting_tokens += len(req.origin_input_ids)
req.time_stats.wait_queue_entry_time = time.perf_counter()
trace_slice_end(RequestStage.REQUEST_PROCESS, req.rid, auto_next_anon=True)
elif self.disaggregation_mode == DisaggregationMode.PREFILL:
@@ -1533,8 +1530,6 @@ class Scheduler(
if abort_existing_req:
self.waiting_queue.pop(idx)
req_to_abort = candidate_req
# update counter for preempted existing request
self.num_waiting_tokens -= len(req_to_abort.origin_input_ids)
message = "The request is aborted by a higher priority request."
self.send_to_tokenizer.send_output(
@@ -1812,15 +1807,9 @@ class Scheduler(
for req in can_run_list:
req.add_latency(RequestStage.PREFILL_WAITING)
can_run_set = set(can_run_list)
waiting_queue = []
for x in self.waiting_queue:
if x in can_run_set:
self.num_waiting_tokens -= len(x.origin_input_ids)
else:
waiting_queue.append(x)
self.waiting_queue = waiting_queue
self.waiting_queue = [
x for x in self.waiting_queue if x not in set(can_run_list)
]
if adder.preempt_list:
for req in adder.preempt_list:
self._add_request_to_queue(req)
@@ -2242,49 +2231,6 @@ class Scheduler(
if_success = False
return if_success
def get_load(self, recv_req: GetLoadReqInput = None) -> GetLoadReqOutput:
if self.is_hybrid:
num_tokens_full = (
self.full_tokens_per_layer
- self.token_to_kv_pool_allocator.full_available_size()
- self.tree_cache.full_evictable_size()
)
num_tokens_swa = (
self.swa_tokens_per_layer
- self.token_to_kv_pool_allocator.swa_available_size()
- self.tree_cache.swa_evictable_size()
)
num_tokens = max(num_tokens_full, num_tokens_swa)
elif self.is_hybrid_gdn:
num_tokens = (
self.max_total_num_tokens
- self.token_to_kv_pool_allocator.available_size()
- self.tree_cache.full_evictable_size()
)
else:
num_tokens = (
self.max_total_num_tokens
- self.token_to_kv_pool_allocator.available_size()
- self.tree_cache.evictable_size()
)
# Tokens in waiting queue, bootstrap queue, prealloc/retract queue
num_tokens += self.num_waiting_tokens
num_waiting_reqs = len(self.waiting_queue)
if self.disaggregation_mode == DisaggregationMode.PREFILL:
num_waiting_reqs += len(self.disagg_prefill_bootstrap_queue.queue)
elif self.disaggregation_mode == DisaggregationMode.DECODE:
num_waiting_reqs += len(self.disagg_decode_prealloc_queue.queue)
num_waiting_reqs += len(self.disagg_decode_transfer_queue.queue)
num_waiting_reqs += len(self.disagg_decode_prealloc_queue.retracted_queue)
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,
)
def get_internal_state(self, recv_req: GetInternalStateReq):
ret = vars(get_global_server_args())
ret["last_gen_throughput"] = self.last_gen_throughput
@@ -2392,9 +2338,6 @@ class Scheduler(
# For disaggregation decode mode, the request in the waiting queue has KV cache allocated.
if self.disaggregation_mode == DisaggregationMode.DECODE:
release_kv_cache(req, self.tree_cache)
else:
# update counter for abort request in prealloc and waiting queue
self.num_waiting_tokens -= len(req.origin_input_ids)
# For mamba radix cache
if req.mamba_pool_idx is not None:

View File

@@ -8,6 +8,7 @@ from typing import TYPE_CHECKING, List, Optional
from sglang.srt.disaggregation.kv_events import EventPublisherFactory, KVEventBatch
from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.environ import envs
from sglang.srt.managers.io_struct import GetLoadReqInput, GetLoadReqOutput
from sglang.srt.managers.schedule_policy import PrefillAdder
from sglang.srt.managers.scheduler import Req, ScheduleBatch
from sglang.srt.metrics.collector import SchedulerMetricsCollector, SchedulerStats
@@ -391,3 +392,31 @@ class SchedulerMetricsMixin:
/ self.stats.max_running_requests_under_SLO,
self.stats.token_usage / 0.9,
)
def get_load(self: Scheduler, _: GetLoadReqInput = None) -> GetLoadReqOutput:
if self.is_hybrid:
full_num_used, swa_num_used, *_ = self._get_swa_token_info()
num_tokens = max(full_num_used, swa_num_used)
elif self.is_hybrid_gdn:
num_tokens = self._get_mamba_token_info()[0]
else:
num_tokens = self._get_token_info()[0]
# Tokens 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,
)