[diffusion] fix: shut down diffusion workers on serve exit (#30110)

This commit is contained in:
Mick
2026-07-04 20:42:39 +08:00
committed by GitHub
parent e552f6ed75
commit 36fc0093d6
4 changed files with 198 additions and 17 deletions
@@ -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):
@@ -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():
@@ -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()
+10 -13
View File
@@ -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: