[diffusion] rl: add sleep/wake support for diffusion engine (#22659)

Co-authored-by: klhhhhh <1412841649@qq.com>
Co-authored-by: zhaochenyang20 <zhaochen20@outlook.com>
Co-authored-by: shuwenn <47200617+alphabetc1@users.noreply.github.com>
Co-authored-by: MikukuOvO <MikukuOvO@users.noreply.github.com>
This commit is contained in:
Fenglin Yu
2026-06-25 09:37:27 +08:00
committed by GitHub
co-authored by klhhhhh zhaochenyang20 shuwenn MikukuOvO
parent 67b2a9ed0c
commit ae6f787637
5 changed files with 329 additions and 0 deletions
@@ -41,6 +41,20 @@ class GetWeightsChecksumReqInput:
module_names: list[str] | None = None
@dataclass
class ReleaseMemoryOccupationReqInput:
"""Request to release (sleep) GPU memory occupation for the diffusion engine."""
pass
@dataclass
class ResumeMemoryOccupationReqInput:
"""Request to resume (wake) GPU memory occupation for the diffusion engine."""
pass
class RolloutRequest(BaseModel):
prompt: str
negative_prompt: Optional[str] = None
@@ -4,6 +4,8 @@ from fastapi import APIRouter, Request
from sglang.multimodal_gen.runtime.entrypoints.post_training.io_struct import (
GetWeightsChecksumReqInput,
ReleaseMemoryOccupationReqInput,
ResumeMemoryOccupationReqInput,
UpdateWeightFromDiskReqInput,
UpdateWeightFromTensorCheckerReqInput,
UpdateWeightFromTensorReqInput,
@@ -39,6 +41,15 @@ async def update_weights_from_disk(request: Request):
status_code=500,
)
if response.output is None:
return orjson_response(
{
"success": False,
"message": response.error or "Unknown status",
},
status_code=500,
)
result = response.output
return orjson_response(
result,
@@ -138,3 +149,51 @@ async def get_weights_checksum(request: Request):
return orjson_response({"error": str(e)}, status_code=500)
return orjson_response(response.output, status_code=200)
@router.post("/release_memory_occupation")
async def release_memory_occupation():
"""Release GPU memory occupation (sleep the engine)."""
try:
response = await async_scheduler_client.forward(
ReleaseMemoryOccupationReqInput()
)
except Exception as e:
return orjson_response({"success": False, "message": str(e)}, status_code=500)
if response.output is None:
return orjson_response(
{
"success": False,
"message": response.error or "Unknown status",
},
status_code=500,
)
payload = response.output
success = bool(payload["success"])
return orjson_response(payload, status_code=200 if success else 400)
@router.post("/resume_memory_occupation")
async def resume_memory_occupation():
"""Resume GPU memory occupation (wake the engine)."""
try:
response = await async_scheduler_client.forward(
ResumeMemoryOccupationReqInput()
)
except Exception as e:
return orjson_response({"success": False, "message": str(e)}, status_code=500)
if response.output is None:
return orjson_response(
{
"success": False,
"message": response.error or "Unknown status",
},
status_code=500,
)
payload = response.output
success = bool(payload["success"])
return orjson_response(payload, status_code=200 if success else 400)
@@ -41,6 +41,9 @@ from sglang.multimodal_gen.runtime.entrypoints.utils import (
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
configure_layerwise_offload_modules,
)
from sglang.multimodal_gen.runtime.managers.memory_managers.memory_occupation_controller import (
MemoryOccupationController,
)
from sglang.multimodal_gen.runtime.pipelines_core import (
ComposedPipelineBase,
LoRAPipeline,
@@ -128,6 +131,7 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
self.cfg_group = get_cfg_group()
self.cfg_cpu_group = self.cfg_group.cpu_group
self._realtime_sessions = RealtimeSessionCache(max_sessions=1)
self.memory_occupation: MemoryOccupationController | None = None
def release_realtime_session(self, session_id: str) -> OutputBatch:
"""release the session of a realtime connection"""
@@ -175,6 +179,18 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
os.environ.get("TRITON_CACHE_DIR"),
)
def is_sleeping(self) -> bool:
return self.memory_occupation.is_sleeping() if self.memory_occupation else False
def _get_memory_occupation(self) -> MemoryOccupationController:
if self.memory_occupation is None:
self.memory_occupation = MemoryOccupationController(
pipeline=self.pipeline,
rank=self.rank,
use_fsdp_inference=self.server_args.use_fsdp_inference,
)
return self.memory_occupation
def init_device_and_model(self) -> None:
"""Initialize the device and load the model."""
torch.get_device_module().set_device(self.local_rank)
@@ -922,6 +938,18 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
status = self.pipeline.get_lora_status()
return OutputBatch(output=status)
def release_memory_occupation(self) -> dict:
return self._get_memory_occupation().release_memory_occupation()
def resume_memory_occupation(self) -> dict:
if self.memory_occupation is None:
return {
"success": True,
"sleeping": False,
"message": "already awake",
}
return self.memory_occupation.resume_memory_occupation()
OOM_MSG = """
OOM detected. Possible solutions:
@@ -0,0 +1,203 @@
# SPDX-License-Identifier: Apache-2.0
import gc
import torch
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
is_layerwise_offloaded_module,
)
from sglang.multimodal_gen.runtime.pipelines_core import ComposedPipelineBase
from sglang.multimodal_gen.runtime.post_training.weights_updater import (
get_updatable_modules,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__)
def _get_module_device(module: torch.nn.Module) -> str:
"""Return best-effort device string for a module."""
param = next(module.parameters(), None)
if param is not None:
return str(param.device)
buffer = next(module.buffers(), None)
if buffer is not None:
return str(buffer.device)
for key, val in vars(module).items():
if key.startswith("_"):
continue
if isinstance(val, torch.Tensor):
return str(val.device)
return "cpu"
def _move_unregistered_tensors(module: torch.nn.Module, device: str) -> None:
"""Move tensor attributes that are not covered by `module.to(device)`."""
def move_tensors(obj):
if torch.is_tensor(obj):
return obj.to(device)
if isinstance(obj, dict):
return {k: move_tensors(v) for k, v in obj.items()}
if isinstance(obj, list):
return [move_tensors(v) for v in obj]
if isinstance(obj, tuple):
return tuple(move_tensors(v) for v in obj)
return obj
attrs = module.__dict__
for attr_name, attr_value in list(attrs.items()):
if attr_name.startswith("_"):
continue
if attr_name in {"_parameters", "_buffers", "_modules"}:
continue
moved_value = move_tensors(attr_value)
if moved_value is not attr_value:
attrs[attr_name] = moved_value
def _is_layerwise_offload_managed(module: torch.nn.Module) -> bool:
return is_layerwise_offloaded_module(module)
class MemoryOccupationController:
def __init__(
self,
pipeline: ComposedPipelineBase | None,
rank: int,
use_fsdp_inference: bool,
):
self.pipeline = pipeline
self.rank = rank
self.use_fsdp_inference = use_fsdp_inference
self._sleeping = False
self._sleep_restore_map: dict[str, str] = {}
def is_sleeping(self) -> bool:
return self._sleeping
def _memory_occupation_result(
self, success: bool, message: str
) -> dict[str, bool | str]:
return {
"success": success,
"sleeping": self._sleeping,
"message": message,
}
@staticmethod
def _clear_torch_device_cache() -> None:
device = torch.get_device_module()
device.synchronize()
gc.collect()
device.empty_cache()
def _move_modules(self, names: list[str], device: str) -> None:
"""
Move selected modules to device.
This function has all-or-nothing semantics:
- Stop on first failure (device query / move / sanitize).
- Roll back modules already moved in this call.
- Raise RuntimeError to caller after rollback.
"""
modules = get_updatable_modules(self.pipeline)
moved: list[str] = []
src_device_map: dict[str, str] = {}
try:
for name in names:
module = modules[name]
src_device_map[name] = _get_module_device(module)
module.to(device)
moved.append(name)
_move_unregistered_tensors(module, device)
except Exception as e:
logger.warning(
f"[_move_modules] move failed, rollback started: target={device} moved={moved} error={e}",
)
for name in moved:
module = modules.get(name)
src_dev = src_device_map.get(name)
module.to(src_dev)
_move_unregistered_tensors(module, src_dev)
raise RuntimeError(
f"failed to move modules to {device}; rollback finished: error={e}"
) from e
def _offload_active_modules_to_cpu(self) -> dict[str, str]:
restore_map: dict[str, str] = {}
for name, module in get_updatable_modules(self.pipeline).items():
if _is_layerwise_offload_managed(module):
continue
device = _get_module_device(module)
if not device.startswith("cpu"):
restore_map[name] = device
self._move_modules(list(restore_map.keys()), "cpu")
self._clear_torch_device_cache()
return restore_map
def _restore_modules_to_original_devices(
self, module_device_map: dict[str, str]
) -> None:
grouped: dict[str, list[str]] = {}
for name, device in module_device_map.items():
grouped.setdefault(device, []).append(name)
for device, names in grouped.items():
self._move_modules(names, device)
def release_memory_occupation(self) -> dict[str, bool | str]:
logger.info(f"[SLEEP] release_memory_occupation rank={self.rank}")
if self._sleeping:
return self._memory_occupation_result(
success=True,
message="already sleeping",
)
if self.use_fsdp_inference:
raise RuntimeError("sleep/wake does not support FSDP inference")
if self.pipeline is None:
return self._memory_occupation_result(
success=False,
message="pipeline not initialized",
)
self._sleep_restore_map = self._offload_active_modules_to_cpu()
self._sleeping = True
return self._memory_occupation_result(
success=True,
message="released GPU memory (moved active modules to CPU)",
)
def resume_memory_occupation(self) -> dict[str, bool | str]:
logger.info(f"[WAKE] resume_memory_occupation rank={self.rank}")
if not self._sleeping:
return self._memory_occupation_result(
success=True,
message="already awake",
)
if self.pipeline is None:
return self._memory_occupation_result(
success=False,
message="pipeline not initialized",
)
if not self._sleep_restore_map:
self._sleeping = False
return self._memory_occupation_result(
success=True,
message="no restore map; marked awake",
)
self._restore_modules_to_original_devices(self._sleep_restore_map)
self._sleep_restore_map = {}
self._sleeping = False
return self._memory_occupation_result(
success=True,
message="resumed GPU memory (restored modules to original devices)",
)
@@ -18,6 +18,8 @@ from sglang.multimodal_gen.runtime.disaggregation.scheduler_mixin import (
)
from sglang.multimodal_gen.runtime.entrypoints.post_training.io_struct import (
GetWeightsChecksumReqInput,
ReleaseMemoryOccupationReqInput,
ResumeMemoryOccupationReqInput,
UpdateWeightFromDiskReqInput,
UpdateWeightFromTensorCheckerReqInput,
UpdateWeightFromTensorReqInput,
@@ -138,6 +140,8 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
self._handle_update_weights_from_tensor_checker
),
GetWeightsChecksumReqInput: self._handle_get_weights_checksum,
ReleaseMemoryOccupationReqInput: self._handle_release_memory_occupation,
ResumeMemoryOccupationReqInput: self._handle_resume_memory_occupation,
}
# FIFO queue entries: (identity, request, enqueue_ts_s)
@@ -214,6 +218,15 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
req = reqs[0]
return self.worker.release_realtime_session(req.session_id)
def _handle_update_weights_from_disk(self, reqs: List[Any]) -> OutputBatch:
"""Handle update_weights_from_disk request for RL workflows."""
if self.worker.is_sleeping():
raise RuntimeError(
"Cannot update weights while the server is sleeping. "
"Call resume_memory_occupation first."
)
return super()._handle_update_weights_from_disk(reqs)
@staticmethod
def _normalize_generation_reqs(reqs: list[Any]) -> list[Req]:
if len(reqs) == 1 and isinstance(reqs[0], list):
@@ -249,6 +262,10 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
):
"""Dispatch generation requests, merging compatible requests when allowed."""
reqs = self._normalize_generation_reqs(reqs)
if self.worker.is_sleeping():
raise RuntimeError(
"Server is sleeping. Call resume_memory_occupation first."
)
warmup_reqs = [req for req in reqs if req.is_warmup]
if warmup_reqs:
self._ensure_warmup_progress_bar(warmup_reqs[0])
@@ -1067,3 +1084,11 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
for pipe in self.result_pipes_from_slaves:
results.append(pipe.recv())
return results
def _handle_release_memory_occupation(self, _reqs: List[Any]) -> OutputBatch:
logger.info(f"[SLEEP] handle_release_memory_occupation on rank={self.gpu_id}")
return OutputBatch(output=self.worker.release_memory_occupation())
def _handle_resume_memory_occupation(self, _reqs: List[Any]) -> OutputBatch:
logger.info(f"[WAKE] handle_resume_memory_occupation on rank={self.gpu_id}")
return OutputBatch(output=self.worker.resume_memory_occupation())