[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:
co-authored by
klhhhhh
zhaochenyang20
shuwenn
MikukuOvO
parent
67b2a9ed0c
commit
ae6f787637
@@ -41,6 +41,20 @@ class GetWeightsChecksumReqInput:
|
|||||||
module_names: list[str] | None = None
|
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):
|
class RolloutRequest(BaseModel):
|
||||||
prompt: str
|
prompt: str
|
||||||
negative_prompt: Optional[str] = None
|
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 (
|
from sglang.multimodal_gen.runtime.entrypoints.post_training.io_struct import (
|
||||||
GetWeightsChecksumReqInput,
|
GetWeightsChecksumReqInput,
|
||||||
|
ReleaseMemoryOccupationReqInput,
|
||||||
|
ResumeMemoryOccupationReqInput,
|
||||||
UpdateWeightFromDiskReqInput,
|
UpdateWeightFromDiskReqInput,
|
||||||
UpdateWeightFromTensorCheckerReqInput,
|
UpdateWeightFromTensorCheckerReqInput,
|
||||||
UpdateWeightFromTensorReqInput,
|
UpdateWeightFromTensorReqInput,
|
||||||
@@ -39,6 +41,15 @@ async def update_weights_from_disk(request: Request):
|
|||||||
status_code=500,
|
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
|
result = response.output
|
||||||
return orjson_response(
|
return orjson_response(
|
||||||
result,
|
result,
|
||||||
@@ -138,3 +149,51 @@ async def get_weights_checksum(request: Request):
|
|||||||
return orjson_response({"error": str(e)}, status_code=500)
|
return orjson_response({"error": str(e)}, status_code=500)
|
||||||
|
|
||||||
return orjson_response(response.output, status_code=200)
|
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 (
|
from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import (
|
||||||
configure_layerwise_offload_modules,
|
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 (
|
from sglang.multimodal_gen.runtime.pipelines_core import (
|
||||||
ComposedPipelineBase,
|
ComposedPipelineBase,
|
||||||
LoRAPipeline,
|
LoRAPipeline,
|
||||||
@@ -128,6 +131,7 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
|
|||||||
self.cfg_group = get_cfg_group()
|
self.cfg_group = get_cfg_group()
|
||||||
self.cfg_cpu_group = self.cfg_group.cpu_group
|
self.cfg_cpu_group = self.cfg_group.cpu_group
|
||||||
self._realtime_sessions = RealtimeSessionCache(max_sessions=1)
|
self._realtime_sessions = RealtimeSessionCache(max_sessions=1)
|
||||||
|
self.memory_occupation: MemoryOccupationController | None = None
|
||||||
|
|
||||||
def release_realtime_session(self, session_id: str) -> OutputBatch:
|
def release_realtime_session(self, session_id: str) -> OutputBatch:
|
||||||
"""release the session of a realtime connection"""
|
"""release the session of a realtime connection"""
|
||||||
@@ -175,6 +179,18 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
|
|||||||
os.environ.get("TRITON_CACHE_DIR"),
|
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:
|
def init_device_and_model(self) -> None:
|
||||||
"""Initialize the device and load the model."""
|
"""Initialize the device and load the model."""
|
||||||
torch.get_device_module().set_device(self.local_rank)
|
torch.get_device_module().set_device(self.local_rank)
|
||||||
@@ -922,6 +938,18 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
|
|||||||
status = self.pipeline.get_lora_status()
|
status = self.pipeline.get_lora_status()
|
||||||
return OutputBatch(output=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_MSG = """
|
||||||
OOM detected. Possible solutions:
|
OOM detected. Possible solutions:
|
||||||
|
|||||||
+203
@@ -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 (
|
from sglang.multimodal_gen.runtime.entrypoints.post_training.io_struct import (
|
||||||
GetWeightsChecksumReqInput,
|
GetWeightsChecksumReqInput,
|
||||||
|
ReleaseMemoryOccupationReqInput,
|
||||||
|
ResumeMemoryOccupationReqInput,
|
||||||
UpdateWeightFromDiskReqInput,
|
UpdateWeightFromDiskReqInput,
|
||||||
UpdateWeightFromTensorCheckerReqInput,
|
UpdateWeightFromTensorCheckerReqInput,
|
||||||
UpdateWeightFromTensorReqInput,
|
UpdateWeightFromTensorReqInput,
|
||||||
@@ -138,6 +140,8 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
|
|||||||
self._handle_update_weights_from_tensor_checker
|
self._handle_update_weights_from_tensor_checker
|
||||||
),
|
),
|
||||||
GetWeightsChecksumReqInput: self._handle_get_weights_checksum,
|
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)
|
# FIFO queue entries: (identity, request, enqueue_ts_s)
|
||||||
@@ -214,6 +218,15 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
|
|||||||
req = reqs[0]
|
req = reqs[0]
|
||||||
return self.worker.release_realtime_session(req.session_id)
|
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
|
@staticmethod
|
||||||
def _normalize_generation_reqs(reqs: list[Any]) -> list[Req]:
|
def _normalize_generation_reqs(reqs: list[Any]) -> list[Req]:
|
||||||
if len(reqs) == 1 and isinstance(reqs[0], list):
|
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."""
|
"""Dispatch generation requests, merging compatible requests when allowed."""
|
||||||
reqs = self._normalize_generation_reqs(reqs)
|
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]
|
warmup_reqs = [req for req in reqs if req.is_warmup]
|
||||||
if warmup_reqs:
|
if warmup_reqs:
|
||||||
self._ensure_warmup_progress_bar(warmup_reqs[0])
|
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:
|
for pipe in self.result_pipes_from_slaves:
|
||||||
results.append(pipe.recv())
|
results.append(pipe.recv())
|
||||||
return results
|
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())
|
||||||
|
|||||||
Reference in New Issue
Block a user