[RL] Handle Mooncake buffers across memory release (#27696)
This commit is contained in:
@@ -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:
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user