diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index c69343149..4e1e24ba6 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -2125,10 +2125,10 @@ 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_used_tokens plus pending tokens not already allocated in the KV pool. + # Used for DP balance. num_total_tokens: int = field( - metadata={"metric": ("gauge", "Used tokens plus pending prefill tokens")} + metadata={"metric": ("gauge", "Used tokens plus pending unallocated tokens")} ) max_total_num_tokens: int = field( metadata={"metric": ("gauge", "Maximum token capacity")} diff --git a/python/sglang/srt/managers/scheduler_components/load_inquirer.py b/python/sglang/srt/managers/scheduler_components/load_inquirer.py index e2b83f6bc..b49409fc7 100644 --- a/python/sglang/srt/managers/scheduler_components/load_inquirer.py +++ b/python/sglang/srt/managers/scheduler_components/load_inquirer.py @@ -104,14 +104,24 @@ class SchedulerLoadInquirer: num_running_reqs = len(self.get_running_batch().reqs) waiting_queues = [self.get_waiting_queue()] + pending_token_queues = [self.get_waiting_queue()] if self.disaggregation_mode == DisaggregationMode.PREFILL: - waiting_queues.append(self.get_disagg_prefill_bootstrap_queue().queue) + prefill_bootstrap_queue = self.get_disagg_prefill_bootstrap_queue().queue + waiting_queues.append(prefill_bootstrap_queue) + pending_token_queues.append(prefill_bootstrap_queue) elif self.disaggregation_mode == DisaggregationMode.DECODE: - waiting_queues.append(self.get_disagg_decode_prealloc_queue().queue) - waiting_queues.append(self.get_disagg_decode_transfer_queue().queue) - waiting_queues.append( + decode_prealloc_queue = self.get_disagg_decode_prealloc_queue().queue + decode_transfer_queue = self.get_disagg_decode_transfer_queue().queue + decode_retracted_queue = ( self.get_disagg_decode_prealloc_queue().retracted_queue ) + waiting_queues.append(decode_prealloc_queue) + waiting_queues.append(decode_transfer_queue) + waiting_queues.append(decode_retracted_queue) + # In disaggregated decode, transfer-queue requests and transferred + # waiting-queue requests have already pre-allocated decode-side KV + # slots, so they are already included in num_used_tokens. + pending_token_queues = [decode_prealloc_queue, decode_retracted_queue] num_waiting_reqs = sum(len(queue) for queue in waiting_queues) num_waiting_uncached_tokens = self.get_num_waiting_uncached_tokens() @@ -119,7 +129,7 @@ class SchedulerLoadInquirer: self.pool_stats_observer.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 + req.seqlen for queue in pending_token_queues for req in queue ) memory = None