[diffusion] chore: bound DiffGenerator local cleanup (#28833)

This commit is contained in:
Mick
2026-06-21 22:23:49 +08:00
committed by GitHub
parent a7f31a6e1b
commit 320b231ea6
3 changed files with 155 additions and 12 deletions
@@ -603,7 +603,7 @@ class DiffGenerator:
# sends the shutdown command to the server
if self.local_scheduler_process and self.owns_scheduler_client:
try:
sync_scheduler_client.forward(ShutdownReq())
sync_scheduler_client.forward(ShutdownReq(), timeout_ms=5000)
except Exception:
pass
@@ -615,12 +615,46 @@ class DiffGenerator:
f"Local worker {process.name} did not terminate gracefully, forcing."
)
process.terminate()
process.join(timeout=1)
if process.is_alive():
process.kill()
process.join(timeout=1)
self.local_scheduler_process = None
if self.owns_scheduler_client:
sync_scheduler_client.close()
self.owns_scheduler_client = False
def _force_shutdown_local_processes(self) -> None:
local_scheduler_process = getattr(self, "local_scheduler_process", None)
log = globals().get("logger")
if local_scheduler_process:
for process in local_scheduler_process:
if process.is_alive():
if log is not None:
log.warning(
f"Local worker {process.name} did not terminate gracefully, forcing."
)
process.terminate()
for process in local_scheduler_process:
process.join(timeout=1)
if process.is_alive():
if log is not None:
log.warning(
f"Local worker {process.name} did not terminate after terminate(), killing."
)
process.kill()
process.join(timeout=1)
self.local_scheduler_process = None
if getattr(self, "owns_scheduler_client", False):
try:
client = globals().get("sync_scheduler_client")
if client is not None:
client.close()
finally:
self.owns_scheduler_client = False
def __enter__(self):
return self
@@ -630,15 +664,18 @@ class DiffGenerator:
def __del__(self):
owns_scheduler_client = bool(getattr(self, "owns_scheduler_client", False))
local_scheduler_process = getattr(self, "local_scheduler_process", None)
log = globals().get("logger")
if owns_scheduler_client:
logger.warning(
"Generator was garbage collected without being shut down. "
"Attempting to shut down the local server and client."
)
self.shutdown()
if log is not None:
log.warning(
"Generator was garbage collected without being shut down. "
"Forcing local server and client cleanup."
)
self._force_shutdown_local_processes()
elif local_scheduler_process:
logger.warning(
"Generator was garbage collected without being shut down. "
"Attempting to shut down the local server."
)
self.shutdown()
if log is not None:
log.warning(
"Generator was garbage collected without being shut down. "
"Forcing local server cleanup."
)
self._force_shutdown_local_processes()
@@ -78,8 +78,12 @@ class SchedulerClient:
f"SchedulerClient connected to backend scheduler at {scheduler_endpoint}"
)
def forward(self, batch: Any) -> Any:
def forward(self, batch: Any, timeout_ms: int | None = None) -> Any:
"""Sends a batch or request to the scheduler and waits for the response."""
previous_timeout_ms = None
if timeout_ms is not None:
previous_timeout_ms = self.scheduler_socket.getsockopt(zmq.RCVTIMEO)
self.scheduler_socket.setsockopt(zmq.RCVTIMEO, timeout_ms)
try:
self.scheduler_socket.send_pyobj(batch)
output_batch = self.scheduler_socket.recv_pyobj()
@@ -88,6 +92,9 @@ class SchedulerClient:
except zmq.error.Again:
logger.error("Timeout waiting for response from scheduler.")
raise TimeoutError("Scheduler did not respond in time.")
finally:
if previous_timeout_ms is not None and self.scheduler_socket is not None:
self.scheduler_socket.setsockopt(zmq.RCVTIMEO, previous_timeout_ms)
def ping(self) -> bool:
"""
@@ -0,0 +1,99 @@
import unittest
from types import SimpleNamespace
from unittest.mock import Mock, patch
from sglang.multimodal_gen.runtime.entrypoints import diffusion_generator as dg
from sglang.multimodal_gen.runtime.entrypoints.diffusion_generator import DiffGenerator
from sglang.multimodal_gen.runtime.entrypoints.utils import ShutdownReq
class _FakeProcess:
name = "fake-worker"
def __init__(self):
self.alive = True
self.join_timeouts = []
self.terminated = False
self.killed = False
def join(self, timeout=None):
self.join_timeouts.append(timeout)
def is_alive(self):
return self.alive
def terminate(self):
self.terminated = True
def kill(self):
self.killed = True
self.alive = False
class TestDiffGeneratorShutdown(unittest.TestCase):
def test_shutdown_uses_bounded_scheduler_timeout_and_forces_worker(self):
generator = object.__new__(DiffGenerator)
process = _FakeProcess()
generator.local_scheduler_process = [process]
generator.owns_scheduler_client = True
client = SimpleNamespace(
forward=Mock(side_effect=TimeoutError("blocked")),
close=Mock(),
)
with patch.object(dg, "sync_scheduler_client", client):
generator.shutdown()
request = client.forward.call_args.args[0]
self.assertIsInstance(request, ShutdownReq)
self.assertEqual(client.forward.call_args.kwargs, {"timeout_ms": 5000})
self.assertTrue(process.terminated)
self.assertTrue(process.killed)
self.assertAlmostEqual(process.join_timeouts[0], 10, delta=0.1)
self.assertEqual(process.join_timeouts[1:], [1, 1])
self.assertIsNone(generator.local_scheduler_process)
self.assertFalse(generator.owns_scheduler_client)
client.close.assert_called_once_with()
def test_del_forces_cleanup_without_scheduler_request(self):
generator = object.__new__(DiffGenerator)
process = _FakeProcess()
generator.local_scheduler_process = [process]
generator.owns_scheduler_client = True
client = SimpleNamespace(
forward=Mock(),
close=Mock(),
)
with patch.object(dg, "sync_scheduler_client", client):
generator.__del__()
client.forward.assert_not_called()
client.close.assert_called_once_with()
self.assertTrue(process.terminated)
self.assertTrue(process.killed)
self.assertIsNone(generator.local_scheduler_process)
self.assertFalse(generator.owns_scheduler_client)
def test_del_tolerates_missing_shutdown_globals(self):
generator = object.__new__(DiffGenerator)
process = _FakeProcess()
generator.local_scheduler_process = [process]
generator.owns_scheduler_client = True
with (
patch.object(dg, "logger", None),
patch.object(dg, "sync_scheduler_client", None),
):
generator.__del__()
self.assertTrue(process.terminated)
self.assertTrue(process.killed)
self.assertIsNone(generator.local_scheduler_process)
self.assertFalse(generator.owns_scheduler_client)
if __name__ == "__main__":
unittest.main()