From 9788c8e867954fdc3a35552ed4d1bbaa1c51672e Mon Sep 17 00:00:00 2001 From: Zilin Zhu Date: Thu, 11 Jun 2026 10:15:00 +0700 Subject: [PATCH] [RL] Handle Mooncake buffers across memory release (#27696) --- python/sglang/srt/disaggregation/decode.py | 18 +++++++++++ .../srt/disaggregation/mooncake/conn.py | 15 ++++++++++ python/sglang/srt/disaggregation/prefill.py | 9 ++++++ python/sglang/srt/managers/scheduler.py | 1 + .../scheduler_components/weight_updater.py | 30 +++++++++++++++++++ 5 files changed, 73 insertions(+) diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index e9efdcdd9..1c0825c86 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -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() diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index b21aee9f7..e2b7e0a9f 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -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) # ------------------------------------------------------------------ diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index ce1afdac3..992bc46d6 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.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: """ diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index af7fb84a0..ac2e76478 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -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, ) diff --git a/python/sglang/srt/managers/scheduler_components/weight_updater.py b/python/sglang/srt/managers/scheduler_components/weight_updater.py index 77bf823b0..c2dc1c1d8 100644 --- a/python/sglang/srt/managers/scheduler_components/weight_updater.py +++ b/python/sglang/srt/managers/scheduler_components/weight_updater.py @@ -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()