[RL] Handle Mooncake buffers across memory release (#27696)
This commit is contained in:
@@ -565,6 +565,16 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
for req in reqs:
|
||||
self.add(req, is_retracted=is_retracted)
|
||||
|
||||
def release_memory_occupation(self):
|
||||
self.queue.clear()
|
||||
self.retracted_queue.clear()
|
||||
if hasattr(self.kv_manager, "deregister_buffer_to_engine"):
|
||||
self.kv_manager.deregister_buffer_to_engine()
|
||||
|
||||
def resume_memory_occupation(self):
|
||||
if hasattr(self.kv_manager, "register_buffer_to_engine"):
|
||||
self.kv_manager.register_buffer_to_engine()
|
||||
|
||||
def resume_retracted_reqs(
|
||||
self, rids_to_check: Optional[List[str]] = None
|
||||
) -> List[Req]:
|
||||
@@ -1698,6 +1708,14 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
|
||||
|
||||
return transferred_reqs
|
||||
|
||||
def release_memory_occupation(self):
|
||||
"""Clean up in-flight transfers before releasing GPU memory."""
|
||||
self.queue.clear()
|
||||
|
||||
def resume_memory_occupation(self):
|
||||
"""Queues are already cleared on release; new transfers can be accepted."""
|
||||
pass
|
||||
|
||||
|
||||
class SchedulerDisaggregationDecodeMixin:
|
||||
@torch.no_grad()
|
||||
|
||||
@@ -257,6 +257,21 @@ class MooncakeKVManager(CommonKVManager):
|
||||
if ptrs and lens:
|
||||
self.engine.batch_register(ptrs, lens)
|
||||
|
||||
def deregister_buffer_to_engine(self):
|
||||
if self.kv_args.kv_data_ptrs:
|
||||
self.engine.batch_deregister(self.kv_args.kv_data_ptrs)
|
||||
|
||||
if self.kv_args.aux_data_ptrs:
|
||||
self.engine.batch_deregister(self.kv_args.aux_data_ptrs)
|
||||
|
||||
for ptrs in self.kv_args.state_data_ptrs or []:
|
||||
if ptrs:
|
||||
self.engine.batch_deregister(ptrs)
|
||||
|
||||
if hasattr(self, "connection_pool"):
|
||||
with self.connection_lock:
|
||||
self.connection_pool.clear()
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Staging buffer methods (all delegate to staging_handler.py)
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@@ -378,6 +378,15 @@ class PrefillBootstrapQueue:
|
||||
else:
|
||||
return bootstrapped_reqs, failed_reqs
|
||||
|
||||
def release_memory_occupation(self):
|
||||
self.queue.clear()
|
||||
if hasattr(self.kv_manager, "deregister_buffer_to_engine"):
|
||||
self.kv_manager.deregister_buffer_to_engine()
|
||||
|
||||
def resume_memory_occupation(self):
|
||||
if hasattr(self.kv_manager, "register_buffer_to_engine"):
|
||||
self.kv_manager.register_buffer_to_engine()
|
||||
|
||||
|
||||
class SchedulerDisaggregationPrefillMixin:
|
||||
"""
|
||||
|
||||
@@ -1587,6 +1587,7 @@ class Scheduler(
|
||||
memory_saver_adapter=self.memory_saver_adapter,
|
||||
flush_cache=self.flush_cache,
|
||||
is_fully_idle=self.is_fully_idle,
|
||||
scheduler=self,
|
||||
metrics_collector=self.metrics_collector,
|
||||
)
|
||||
|
||||
|
||||
@@ -16,6 +16,7 @@ from sglang.srt.constants import (
|
||||
GPU_MEMORY_TYPE_KV_CACHE,
|
||||
GPU_MEMORY_TYPE_WEIGHTS,
|
||||
)
|
||||
from sglang.srt.disaggregation.utils import DisaggregationMode
|
||||
from sglang.srt.managers.io_struct import (
|
||||
CheckWeightsReqInput,
|
||||
CheckWeightsReqOutput,
|
||||
@@ -77,6 +78,7 @@ class SchedulerWeightUpdaterManager:
|
||||
memory_saver_adapter: Any
|
||||
flush_cache: Callable[..., bool]
|
||||
is_fully_idle: Callable[..., bool]
|
||||
scheduler: Optional[Any] = None
|
||||
metrics_collector: Optional[Any] = None
|
||||
offload_tags: set = field(default_factory=set)
|
||||
stashed_model_static_state: Any = None
|
||||
@@ -189,6 +191,20 @@ class SchedulerWeightUpdaterManager:
|
||||
self.offload_tags.add(tag)
|
||||
|
||||
if GPU_MEMORY_TYPE_KV_CACHE in tags:
|
||||
scheduler = self.scheduler
|
||||
if scheduler is not None:
|
||||
if scheduler.disaggregation_mode == DisaggregationMode.DECODE:
|
||||
for queue_name in (
|
||||
"disagg_decode_transfer_queue",
|
||||
"disagg_decode_prealloc_queue",
|
||||
):
|
||||
queue = getattr(scheduler, queue_name, None)
|
||||
if queue is not None:
|
||||
queue.release_memory_occupation()
|
||||
elif scheduler.disaggregation_mode == DisaggregationMode.PREFILL:
|
||||
queue = getattr(scheduler, "disagg_prefill_bootstrap_queue", None)
|
||||
if queue is not None:
|
||||
queue.release_memory_occupation()
|
||||
self.memory_saver_adapter.pause(GPU_MEMORY_TYPE_KV_CACHE)
|
||||
self.flush_cache()
|
||||
|
||||
@@ -229,6 +245,20 @@ class SchedulerWeightUpdaterManager:
|
||||
|
||||
if GPU_MEMORY_TYPE_KV_CACHE in tags:
|
||||
self.memory_saver_adapter.resume(GPU_MEMORY_TYPE_KV_CACHE)
|
||||
scheduler = self.scheduler
|
||||
if scheduler is not None:
|
||||
if scheduler.disaggregation_mode == DisaggregationMode.DECODE:
|
||||
for queue_name in (
|
||||
"disagg_decode_transfer_queue",
|
||||
"disagg_decode_prealloc_queue",
|
||||
):
|
||||
queue = getattr(scheduler, queue_name, None)
|
||||
if queue is not None:
|
||||
queue.resume_memory_occupation()
|
||||
elif scheduler.disaggregation_mode == DisaggregationMode.PREFILL:
|
||||
queue = getattr(scheduler, "disagg_prefill_bootstrap_queue", None)
|
||||
if queue is not None:
|
||||
queue.resume_memory_occupation()
|
||||
|
||||
return ResumeMemoryOccupationReqOutput()
|
||||
|
||||
|
||||
Reference in New Issue
Block a user