CUDA-graph-compatible releasing and resuming KV cache and model weight memory (#2630)

This commit is contained in:
fzyzcjy
2025-01-13 11:38:51 -08:00
committed by GitHub
parent d08c77c434
commit 923f518337
12 changed files with 406 additions and 60 deletions
@@ -53,6 +53,10 @@ from sglang.srt.managers.io_struct import (
OpenSessionReqInput,
OpenSessionReqOutput,
ProfileReq,
ReleaseMemoryOccupationReqInput,
ReleaseMemoryOccupationReqOutput,
ResumeMemoryOccupationReqInput,
ResumeMemoryOccupationReqOutput,
SessionParams,
TokenizedEmbeddingReqInput,
TokenizedGenerateReqInput,
@@ -188,6 +192,12 @@ class TokenizerManager:
self.get_weights_by_name_communicator = _Communicator(
self.send_to_scheduler, server_args.dp_size
)
self.release_memory_occupation_communicator = _Communicator(
self.send_to_scheduler, server_args.dp_size
)
self.resume_memory_occupation_communicator = _Communicator(
self.send_to_scheduler, server_args.dp_size
)
# Metrics
if self.enable_metrics:
@@ -548,6 +558,22 @@ class TokenizerManager:
else:
return all_parameters
async def release_memory_occupation(
self,
obj: ReleaseMemoryOccupationReqInput,
request: Optional[fastapi.Request] = None,
):
self.auto_create_handle_loop()
await self.release_memory_occupation_communicator(obj)
async def resume_memory_occupation(
self,
obj: ResumeMemoryOccupationReqInput,
request: Optional[fastapi.Request] = None,
):
self.auto_create_handle_loop()
await self.resume_memory_occupation_communicator(obj)
async def open_session(
self, obj: OpenSessionReqInput, request: Optional[fastapi.Request] = None
):
@@ -627,6 +653,8 @@ class TokenizerManager:
UpdateWeightsFromDistributedReqOutput,
GetWeightsByNameReqOutput,
InitWeightsUpdateGroupReqOutput,
ReleaseMemoryOccupationReqOutput,
ResumeMemoryOccupationReqOutput,
] = await self.recv_from_detokenizer.recv_pyobj()
if isinstance(recv_obj, (BatchStrOut, BatchEmbeddingOut, BatchTokenIDOut)):
@@ -709,6 +737,10 @@ class TokenizerManager:
self.update_weights_from_tensor_communicator.handle_recv(recv_obj)
elif isinstance(recv_obj, GetWeightsByNameReqOutput):
self.get_weights_by_name_communicator.handle_recv(recv_obj)
elif isinstance(recv_obj, ReleaseMemoryOccupationReqOutput):
self.release_memory_occupation_communicator.handle_recv(recv_obj)
elif isinstance(recv_obj, ResumeMemoryOccupationReqOutput):
self.resume_memory_occupation_communicator.handle_recv(recv_obj)
else:
raise ValueError(f"Invalid object: {recv_obj=}")