Fix disaggregated decode load token accounting (#25736)

This commit is contained in:
weireweire
2026-06-15 11:37:41 +08:00
committed by GitHub
parent 0417951a86
commit bf38a0b03d
2 changed files with 18 additions and 8 deletions
+3 -3
View File
@@ -2125,10 +2125,10 @@ class GetLoadsReqOutput(BaseReq):
num_used_tokens: int = field( num_used_tokens: int = field(
metadata={"metric": ("gauge", "Number of tokens in use")} metadata={"metric": ("gauge", "Number of tokens in use")}
) )
# num_used_tokens + pending prefill tokens (waiting-queue seqlen, incl. # num_used_tokens plus pending tokens not already allocated in the KV pool.
# disagg bootstrap/prealloc/transfer queues). Used for DP balance. # Used for DP balance.
num_total_tokens: int = field( 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( max_total_num_tokens: int = field(
metadata={"metric": ("gauge", "Maximum token capacity")} metadata={"metric": ("gauge", "Maximum token capacity")}
@@ -104,14 +104,24 @@ class SchedulerLoadInquirer:
num_running_reqs = len(self.get_running_batch().reqs) num_running_reqs = len(self.get_running_batch().reqs)
waiting_queues = [self.get_waiting_queue()] waiting_queues = [self.get_waiting_queue()]
pending_token_queues = [self.get_waiting_queue()]
if self.disaggregation_mode == DisaggregationMode.PREFILL: 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: elif self.disaggregation_mode == DisaggregationMode.DECODE:
waiting_queues.append(self.get_disagg_decode_prealloc_queue().queue) decode_prealloc_queue = self.get_disagg_decode_prealloc_queue().queue
waiting_queues.append(self.get_disagg_decode_transfer_queue().queue) decode_transfer_queue = self.get_disagg_decode_transfer_queue().queue
waiting_queues.append( decode_retracted_queue = (
self.get_disagg_decode_prealloc_queue().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_reqs = sum(len(queue) for queue in waiting_queues)
num_waiting_uncached_tokens = self.get_num_waiting_uncached_tokens() 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() self.pool_stats_observer.get_pool_stats().get_kv_token_stats()
) )
num_total_tokens = num_used_tokens + sum( 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 memory = None