diff --git a/python/sglang/multimodal_gen/runtime/launch_server.py b/python/sglang/multimodal_gen/runtime/launch_server.py index fcb8f2f1e..0e4ab4938 100644 --- a/python/sglang/multimodal_gen/runtime/launch_server.py +++ b/python/sglang/multimodal_gen/runtime/launch_server.py @@ -6,6 +6,7 @@ import os import signal import sys import threading +import time import psutil import uvicorn @@ -15,7 +16,9 @@ from sglang.multimodal_gen.runtime.disaggregation.orchestrator import ( ) from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType from sglang.multimodal_gen.runtime.entrypoints.http_server import create_app +from sglang.multimodal_gen.runtime.entrypoints.utils import ShutdownReq from sglang.multimodal_gen.runtime.managers.gpu_worker import run_scheduler_process +from sglang.multimodal_gen.runtime.scheduler_client import SchedulerClient from sglang.multimodal_gen.runtime.server_args import ( ServerArgs, prepare_server_args, @@ -24,6 +27,12 @@ from sglang.multimodal_gen.runtime.server_args import ( from sglang.multimodal_gen.runtime.utils.common import is_port_available from sglang.multimodal_gen.runtime.utils.logging_utils import configure_logger, logger from sglang.multimodal_gen.runtime.utils.trace_wrapper import init_diffusion_tracing +from sglang.multimodal_gen.utils import kill_itself_when_parent_died + +_SCHEDULER_SHUTDOWN_TIMEOUT_MS = 5000 +_WORKER_JOIN_TIMEOUT_S = 10 +_WORKER_TERMINATE_TIMEOUT_S = 1 +_WORKER_KILL_TIMEOUT_S = 1 def _find_available_port( @@ -83,6 +92,82 @@ def kill_process_tree(parent_pid, include_parent: bool = True, skip_pid: int = N pass +def _process_names(processes) -> str: + return ", ".join(getattr(p, "name", repr(p)) for p in processes) + + +def _join_processes_with_deadline(processes, timeout_s: float) -> None: + deadline = time.monotonic() + timeout_s + for process in processes: + remaining_s = max(0.0, deadline - time.monotonic()) + process.join(timeout=remaining_s) + + +def _terminate_alive_processes(processes, timeout_s: float) -> list: + alive = [p for p in processes if p.is_alive()] + if not alive: + return [] + + logger.warning( + "Worker process(es) did not exit in time; terminating: %s", + _process_names(alive), + ) + for process in alive: + process.terminate() + _join_processes_with_deadline(alive, timeout_s) + return [p for p in alive if p.is_alive()] + + +def _kill_alive_processes(processes, timeout_s: float) -> None: + alive = [p for p in processes if p.is_alive()] + if not alive: + return + + logger.warning( + "Worker process(es) did not terminate in time; killing: %s", + _process_names(alive), + ) + for process in alive: + process.kill() + _join_processes_with_deadline(alive, timeout_s) + + +def _run_http_server_process(server_args: ServerArgs) -> None: + kill_itself_when_parent_died() + launch_http_server_only(server_args) + + +def _request_monolithic_scheduler_shutdown(server_args: ServerArgs) -> None: + if server_args.disagg_role != RoleType.MONOLITHIC: + return + + client = SchedulerClient() + try: + client.initialize(server_args) + client.forward(ShutdownReq(), timeout_ms=_SCHEDULER_SHUTDOWN_TIMEOUT_MS) + except Exception as e: + logger.warning("Failed to request graceful scheduler shutdown: %s", e) + finally: + client.close() + + +def shutdown_scheduler_processes( + server_args: ServerArgs | None, + processes: list, + *, + request_shutdown: bool = True, +) -> None: + if not processes: + return + + if request_shutdown and server_args is not None: + _request_monolithic_scheduler_shutdown(server_args) + + _join_processes_with_deadline(processes, _WORKER_JOIN_TIMEOUT_S) + alive = _terminate_alive_processes(processes, _WORKER_TERMINATE_TIMEOUT_S) + _kill_alive_processes(alive, _WORKER_KILL_TIMEOUT_S) + + def launch_server(server_args: ServerArgs, launch_http_server: bool = True): """ Args: @@ -198,14 +283,17 @@ def launch_server(server_args: ServerArgs, launch_http_server: bool = True): if server_args.webui: logger.info("Launch FastAPI server in another process because of webui.") http_server_process = mp.Process( - target=launch_http_server_only, + target=_run_http_server_process, args=(server_args,), name="sglang-diffusion-webui", daemon=True, ) http_server_process.start() else: - launch_http_server_only(server_args) + try: + launch_http_server_only(server_args) + finally: + shutdown_scheduler_processes(server_args, processes) return processes @@ -408,7 +496,13 @@ def launch_pool_disagg_server( "Starting FastAPI server (connected to DiffusionServer at port %d).", server_args.scheduler_port, ) - launch_http_server_only(server_args) + try: + launch_http_server_only(server_args) + finally: + diffusion_server.stop() + shutdown_scheduler_processes( + server_args, all_processes, request_shutdown=False + ) return all_processes @@ -538,7 +632,10 @@ def launch_disagg_server(server_args: ServerArgs): "Starting HTTP server (connected to DiffusionServer at port %d).", base_port, ) - launch_http_server_only(server_args) + try: + launch_http_server_only(server_args) + finally: + diffusion_server.stop() def launch_disagg_role(server_args: ServerArgs): @@ -659,6 +756,8 @@ def launch_disagg_role(server_args: ServerArgs): p.join() except KeyboardInterrupt: logger.info("Role %s shutting down.", role_type.value) + finally: + shutdown_scheduler_processes(role_args, processes, request_shutdown=False) def dispatch_launch(server_args: ServerArgs): diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py index b3aad006d..3f943c92a 100644 --- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py +++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py @@ -78,6 +78,7 @@ from sglang.multimodal_gen.runtime.utils.trace_wrapper import ( init_diffusion_tracing, trace_slice, ) +from sglang.multimodal_gen.utils import kill_itself_when_parent_died from sglang.srt.utils.network import NetworkAddress logger = init_logger(__name__) @@ -999,6 +1000,7 @@ def run_scheduler_process( Rank 0 acts as the master, handling ZMQ requests and coordinating slaves. Ranks > 0 act as slaves, waiting for tasks from the master. """ + kill_itself_when_parent_died() configure_logger(server_args) globally_suppress_loggers() if current_platform.is_cuda(): diff --git a/python/sglang/multimodal_gen/test/unit/test_launch_server_shutdown.py b/python/sglang/multimodal_gen/test/unit/test_launch_server_shutdown.py new file mode 100644 index 000000000..308ce5b50 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_launch_server_shutdown.py @@ -0,0 +1,83 @@ +import unittest +from types import SimpleNamespace +from unittest.mock import Mock, patch + +from sglang.multimodal_gen.runtime import launch_server as ls +from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType +from sglang.multimodal_gen.runtime.entrypoints.utils import ShutdownReq + + +class _FakeProcess: + name = "fake-worker" + + def __init__(self, *, exit_on_join: bool = False): + self.alive = True + self.exit_on_join = exit_on_join + self.join_timeouts = [] + self.terminated = False + self.killed = False + + def join(self, timeout=None): + self.join_timeouts.append(timeout) + if self.exit_on_join: + self.alive = False + + def is_alive(self): + return self.alive + + def terminate(self): + self.terminated = True + + def kill(self): + self.killed = True + self.alive = False + + +class TestLaunchServerShutdown(unittest.TestCase): + def test_monolithic_shutdown_requests_scheduler_then_forces_worker(self): + process = _FakeProcess() + server_args = SimpleNamespace(disagg_role=RoleType.MONOLITHIC) + client = Mock() + + with patch.object(ls, "SchedulerClient", return_value=client): + ls.shutdown_scheduler_processes(server_args, [process]) + + client.initialize.assert_called_once_with(server_args) + request = client.forward.call_args.args[0] + self.assertIsInstance(request, ShutdownReq) + self.assertEqual(client.forward.call_args.kwargs, {"timeout_ms": 5000}) + client.close.assert_called_once_with() + + self.assertTrue(process.terminated) + self.assertTrue(process.killed) + self.assertAlmostEqual(process.join_timeouts[0], 10, delta=0.1) + self.assertAlmostEqual(process.join_timeouts[1], 1, delta=0.1) + self.assertAlmostEqual(process.join_timeouts[2], 1, delta=0.1) + + def test_scheduler_shutdown_error_still_forces_worker(self): + process = _FakeProcess() + server_args = SimpleNamespace(disagg_role=RoleType.MONOLITHIC) + client = Mock() + client.forward.side_effect = TimeoutError("blocked") + + with patch.object(ls, "SchedulerClient", return_value=client): + ls.shutdown_scheduler_processes(server_args, [process]) + + client.close.assert_called_once_with() + self.assertTrue(process.terminated) + self.assertTrue(process.killed) + + def test_disagg_role_does_not_send_monolithic_shutdown_request(self): + process = _FakeProcess(exit_on_join=True) + server_args = SimpleNamespace(disagg_role=RoleType.ENCODER) + + with patch.object(ls, "SchedulerClient") as scheduler_client: + ls.shutdown_scheduler_processes(server_args, [process]) + + scheduler_client.assert_not_called() + self.assertFalse(process.terminated) + self.assertFalse(process.killed) + + +if __name__ == "__main__": + unittest.main() diff --git a/python/sglang/multimodal_gen/utils.py b/python/sglang/multimodal_gen/utils.py index 359fd35f3..2f248451c 100644 --- a/python/sglang/multimodal_gen/utils.py +++ b/python/sglang/multimodal_gen/utils.py @@ -527,21 +527,18 @@ def shallow_asdict(obj) -> dict[str, Any]: return {f.name: getattr(obj, f.name) for f in fields(obj)} -# TODO: validate that this is fine def kill_itself_when_parent_died() -> None: - # if sys.platform == "linux": - # sigkill this process when parent worker manager dies - PR_SET_PDEATHSIG = 1 - import platform + if sys.platform != "linux": + return - if platform.system() == "Linux": - libc = ctypes.CDLL("libc.so.6") - libc.prctl(PR_SET_PDEATHSIG, signal.SIGKILL) - # elif platform.system() == "Darwin": - # libc = ctypes.CDLL("libc.dylib") - # logger.warning("kill_itself_when_parent_died is only supported in linux.") - else: - logger.warning("kill_itself_when_parent_died is only supported in linux.") + # keep GPU workers tied to the CLI process even if the parent is SIGKILLed + PR_SET_PDEATHSIG = 1 + libc = ctypes.CDLL("libc.so.6", use_errno=True) + if libc.prctl(PR_SET_PDEATHSIG, signal.SIGKILL) != 0: + err = ctypes.get_errno() + raise OSError(err, os.strerror(err)) + if os.getppid() == 1: + os.kill(os.getpid(), signal.SIGKILL) def get_exception_traceback() -> str: