Extract release_req and retract_all as module-level free functions (#26548)
This commit is contained in:
@@ -1492,6 +1492,53 @@ def set_mamba_track_indices_from_reqs(batch):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def release_req(
|
||||||
|
*,
|
||||||
|
req: Req,
|
||||||
|
remaing_req_count: int,
|
||||||
|
server_args: ServerArgs,
|
||||||
|
req_to_token_pool: ReqToTokenPool,
|
||||||
|
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator,
|
||||||
|
tree_cache: BasePrefixCache,
|
||||||
|
hisparse_coordinator: Optional[HiSparseCoordinator],
|
||||||
|
) -> None:
|
||||||
|
if hisparse_coordinator is not None and not req.finished():
|
||||||
|
hisparse_coordinator.retract_req(req)
|
||||||
|
|
||||||
|
if server_args.disaggregation_mode == "decode":
|
||||||
|
req.offload_kv_cache(req_to_token_pool, token_to_kv_pool_allocator)
|
||||||
|
# TODO (csy): for preempted requests, we may want to insert into the tree
|
||||||
|
release_kv_cache(req, tree_cache, is_insert=False)
|
||||||
|
# NOTE(lsyin): we should use the newly evictable memory instantly.
|
||||||
|
num_tokens = remaing_req_count * envs.SGLANG_RETRACT_DECODE_STEPS.get()
|
||||||
|
evict_from_tree_cache(tree_cache, num_tokens)
|
||||||
|
|
||||||
|
req.reset_for_retract()
|
||||||
|
|
||||||
|
|
||||||
|
def retract_all(
|
||||||
|
*,
|
||||||
|
reqs: List[Req],
|
||||||
|
server_args: ServerArgs,
|
||||||
|
req_to_token_pool: ReqToTokenPool,
|
||||||
|
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator,
|
||||||
|
tree_cache: BasePrefixCache,
|
||||||
|
hisparse_coordinator: Optional[HiSparseCoordinator],
|
||||||
|
) -> List[Req]:
|
||||||
|
retracted_reqs = reqs
|
||||||
|
for idx in range(len(reqs)):
|
||||||
|
release_req(
|
||||||
|
req=reqs[idx],
|
||||||
|
remaing_req_count=len(reqs) - idx,
|
||||||
|
server_args=server_args,
|
||||||
|
req_to_token_pool=req_to_token_pool,
|
||||||
|
token_to_kv_pool_allocator=token_to_kv_pool_allocator,
|
||||||
|
tree_cache=tree_cache,
|
||||||
|
hisparse_coordinator=hisparse_coordinator,
|
||||||
|
)
|
||||||
|
return retracted_reqs
|
||||||
|
|
||||||
|
|
||||||
def _compute_chunked_req_next_prompt_token(
|
def _compute_chunked_req_next_prompt_token(
|
||||||
chunked_req: Optional[Req],
|
chunked_req: Optional[Req],
|
||||||
) -> Optional[int]:
|
) -> Optional[int]:
|
||||||
@@ -2292,10 +2339,14 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
return self.token_to_kv_pool_allocator.available_size() >= num_tokens
|
return self.token_to_kv_pool_allocator.available_size() >= num_tokens
|
||||||
|
|
||||||
def retract_all(self, server_args: ServerArgs):
|
def retract_all(self, server_args: ServerArgs):
|
||||||
retracted_reqs = self.reqs
|
retracted_reqs = retract_all(
|
||||||
for idx in range(len(self.reqs)):
|
reqs=self.reqs,
|
||||||
self.release_req(idx, len(self.reqs) - idx, server_args)
|
server_args=server_args,
|
||||||
|
req_to_token_pool=self.req_to_token_pool,
|
||||||
|
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
|
||||||
|
tree_cache=self.tree_cache,
|
||||||
|
hisparse_coordinator=self.hisparse_coordinator,
|
||||||
|
)
|
||||||
self.reqs = []
|
self.reqs = []
|
||||||
return retracted_reqs
|
return retracted_reqs
|
||||||
|
|
||||||
@@ -2364,22 +2415,15 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
return retracted_reqs, new_estimate_ratio, reqs_to_abort
|
return retracted_reqs, new_estimate_ratio, reqs_to_abort
|
||||||
|
|
||||||
def release_req(self, idx: int, remaing_req_count: int, server_args: ServerArgs):
|
def release_req(self, idx: int, remaing_req_count: int, server_args: ServerArgs):
|
||||||
req = self.reqs[idx]
|
release_req(
|
||||||
|
req=self.reqs[idx],
|
||||||
if self.hisparse_coordinator is not None and not req.finished():
|
remaing_req_count=remaing_req_count,
|
||||||
self.hisparse_coordinator.retract_req(req)
|
server_args=server_args,
|
||||||
|
req_to_token_pool=self.req_to_token_pool,
|
||||||
if server_args.disaggregation_mode == "decode":
|
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
|
||||||
req.offload_kv_cache(
|
tree_cache=self.tree_cache,
|
||||||
self.req_to_token_pool, self.token_to_kv_pool_allocator
|
hisparse_coordinator=self.hisparse_coordinator,
|
||||||
)
|
)
|
||||||
# TODO (csy): for preempted requests, we may want to insert into the tree
|
|
||||||
release_kv_cache(req, self.tree_cache, is_insert=False)
|
|
||||||
# NOTE(lsyin): we should use the newly evictable memory instantly.
|
|
||||||
num_tokens = remaing_req_count * envs.SGLANG_RETRACT_DECODE_STEPS.get()
|
|
||||||
evict_from_tree_cache(self.tree_cache, num_tokens)
|
|
||||||
|
|
||||||
req.reset_for_retract()
|
|
||||||
|
|
||||||
def prepare_encoder_info_decode(self):
|
def prepare_encoder_info_decode(self):
|
||||||
# Reset the encoder cached status
|
# Reset the encoder cached status
|
||||||
|
|||||||
Reference in New Issue
Block a user