Fix disaggregated decode load token accounting (#25736)
This commit is contained in:
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user