[RL] Handle Mooncake buffers across memory release (#27696)

This commit is contained in:
Zilin Zhu
2026-06-11 11:15:00 +08:00
committed by GitHub
parent f8b0a120b8
commit 9788c8e867
5 changed files with 73 additions and 0 deletions
@@ -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:
"""
+1
View File
@@ -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()