[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: for req in reqs:
self.add(req, is_retracted=is_retracted) 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( def resume_retracted_reqs(
self, rids_to_check: Optional[List[str]] = None self, rids_to_check: Optional[List[str]] = None
) -> List[Req]: ) -> List[Req]:
@@ -1698,6 +1708,14 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
return transferred_reqs 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: class SchedulerDisaggregationDecodeMixin:
@torch.no_grad() @torch.no_grad()
@@ -257,6 +257,21 @@ class MooncakeKVManager(CommonKVManager):
if ptrs and lens: if ptrs and lens:
self.engine.batch_register(ptrs, 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) # Staging buffer methods (all delegate to staging_handler.py)
# ------------------------------------------------------------------ # ------------------------------------------------------------------
@@ -378,6 +378,15 @@ class PrefillBootstrapQueue:
else: else:
return bootstrapped_reqs, failed_reqs 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: class SchedulerDisaggregationPrefillMixin:
""" """
+1
View File
@@ -1587,6 +1587,7 @@ class Scheduler(
memory_saver_adapter=self.memory_saver_adapter, memory_saver_adapter=self.memory_saver_adapter,
flush_cache=self.flush_cache, flush_cache=self.flush_cache,
is_fully_idle=self.is_fully_idle, is_fully_idle=self.is_fully_idle,
scheduler=self,
metrics_collector=self.metrics_collector, metrics_collector=self.metrics_collector,
) )
@@ -16,6 +16,7 @@ from sglang.srt.constants import (
GPU_MEMORY_TYPE_KV_CACHE, GPU_MEMORY_TYPE_KV_CACHE,
GPU_MEMORY_TYPE_WEIGHTS, GPU_MEMORY_TYPE_WEIGHTS,
) )
from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.managers.io_struct import ( from sglang.srt.managers.io_struct import (
CheckWeightsReqInput, CheckWeightsReqInput,
CheckWeightsReqOutput, CheckWeightsReqOutput,
@@ -77,6 +78,7 @@ class SchedulerWeightUpdaterManager:
memory_saver_adapter: Any memory_saver_adapter: Any
flush_cache: Callable[..., bool] flush_cache: Callable[..., bool]
is_fully_idle: Callable[..., bool] is_fully_idle: Callable[..., bool]
scheduler: Optional[Any] = None
metrics_collector: Optional[Any] = None metrics_collector: Optional[Any] = None
offload_tags: set = field(default_factory=set) offload_tags: set = field(default_factory=set)
stashed_model_static_state: Any = None stashed_model_static_state: Any = None
@@ -189,6 +191,20 @@ class SchedulerWeightUpdaterManager:
self.offload_tags.add(tag) self.offload_tags.add(tag)
if GPU_MEMORY_TYPE_KV_CACHE in tags: 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.memory_saver_adapter.pause(GPU_MEMORY_TYPE_KV_CACHE)
self.flush_cache() self.flush_cache()
@@ -229,6 +245,20 @@ class SchedulerWeightUpdaterManager:
if GPU_MEMORY_TYPE_KV_CACHE in tags: if GPU_MEMORY_TYPE_KV_CACHE in tags:
self.memory_saver_adapter.resume(GPU_MEMORY_TYPE_KV_CACHE) 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() return ResumeMemoryOccupationReqOutput()