[Fix] Release NCCL on scheduler exit and let the ASGI server own shutdown (#39560)

This commit is contained in:
Liangsheng Yin
2026-09-15 13:39:15 -07:00
committed by GitHub
parent 3c48c1e967
commit d58342deab
5 changed files with 111 additions and 19 deletions
@@ -3123,6 +3123,26 @@ def destroy_distributed_environment():
torch.distributed.destroy_process_group()
def abort_distributed_environment() -> None:
"""Drop this rank's communicators locally.
``destroy_process_group`` is collective and blocks when a peer is gone,
which on a shutdown path is the common case.
"""
if not torch.distributed.is_initialized():
return
abort = getattr(torch.distributed.distributed_c10d, "_abort_process_group", None)
if abort is None:
# Older torch exposes no non-collective teardown,
# and the collective one is what this function exists to avoid.
return
try:
# No argument aborts every group, the default one included.
abort()
except Exception as e:
logger.warning(f"NCCL abort on shutdown failed, {type(e).__name__}: {e}")
def cleanup_dist_env_and_memory(shutdown_ray: bool = False):
destroy_model_parallel()
destroy_distributed_environment()
+32 -15
View File
@@ -2457,6 +2457,7 @@ def _run_granian_server(
log_level,
http2_max_concurrent_streams,
http2_initial_connection_window_size,
tokenizer_manager=None,
tokenizer_worker_num=1,
ssl_certfile=None,
ssl_keyfile=None,
@@ -2514,6 +2515,10 @@ def _run_granian_server(
server = Server(**granian_kwargs)
if tokenizer_worker_num == 1:
if tokenizer_manager is not None:
# auto_create_handle_loop replaces the signal handler wired below,
# so shutdown can only reach this server through the hook.
tokenizer_manager.set_server_stop_hook(server.stop)
async def serve():
# The embedded server does not install its own signal handlers, so wire
@@ -2638,6 +2643,7 @@ def _setup_and_run_http_server(
ssl_ca_certs=get_serving().ssl_ca_certs,
ssl_keyfile_password=get_serving().ssl_keyfile_password,
ssl_verify=False, # No MTLS supported for now.
tokenizer_manager=tokenizer_manager,
)
elif get_serving().enable_ssl_refresh:
# Use Config/Server API for access to the SSLContext.
@@ -2660,6 +2666,9 @@ def _setup_and_run_http_server(
from sglang.srt.entrypoints.ssl_utils import SSLCertRefresher
server = uvicorn.Server(config)
tokenizer_manager.set_server_stop_hook(
lambda: setattr(server, "should_exit", True)
)
async def _run_with_ssl_refresh():
refresher = SSLCertRefresher(
@@ -2678,23 +2687,31 @@ def _setup_and_run_http_server(
asyncio.run(_run_with_ssl_refresh())
else:
# Default case, one tokenizer process
uvicorn.run(
app,
host=get_serving().host,
port=get_serving().port,
root_path=get_serving().fastapi_root_path,
log_level=get_observability().log_level_http
or get_observability().log_level,
timeout_keep_alive=envs.SGLANG_TIMEOUT_KEEP_ALIVE.get(),
loop="uvloop",
ssl_keyfile=get_serving().ssl_keyfile,
ssl_certfile=get_serving().ssl_certfile,
ssl_ca_certs=get_serving().ssl_ca_certs,
ssl_keyfile_password=get_serving().ssl_keyfile_password,
# Default case, one tokenizer process.
# A Server rather than uvicorn.run(), so shutdown can ask it to stop.
server = uvicorn.Server(
uvicorn.Config(
app,
host=get_serving().host,
port=get_serving().port,
root_path=get_serving().fastapi_root_path,
log_level=get_observability().log_level_http
or get_observability().log_level,
timeout_keep_alive=envs.SGLANG_TIMEOUT_KEEP_ALIVE.get(),
loop="uvloop",
ssl_keyfile=get_serving().ssl_keyfile,
ssl_certfile=get_serving().ssl_certfile,
ssl_ca_certs=get_serving().ssl_ca_certs,
ssl_keyfile_password=get_serving().ssl_keyfile_password,
)
)
tokenizer_manager.set_server_stop_hook(
lambda: setattr(server, "should_exit", True)
)
server.run()
else:
# Multiple tokenizer and http processes
# Multiple tokenizer and http processes.
# Child processes re-import the app, so no stop hook here.
from uvicorn.config import LOGGING_CONFIG
LOGGING_CONFIG["loggers"]["sglang.srt.entrypoints.http_server"] = {
+6 -1
View File
@@ -98,7 +98,10 @@ from sglang.srt.disaggregation.utils import (
unified_memory_disagg_move_gate,
)
from sglang.srt.distributed import get_pp_group, get_world_group
from sglang.srt.distributed.parallel_state import get_tp_group
from sglang.srt.distributed.parallel_state import (
abort_distributed_environment,
get_tp_group,
)
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.dllm.mixin.scheduler import SchedulerDllmMixin
from sglang.srt.environ import envs, exportable_env_vars
@@ -5873,6 +5876,8 @@ def run_scheduler_process(
# and the synchronize() in destroy() could itself hang.
if scheduler.gracefully_exit:
scheduler.release_host_resources()
# Last: anything above may still need a working communicator.
abort_distributed_environment()
def _make_abort_req(
@@ -34,7 +34,17 @@ from datetime import datetime
from enum import Enum
from functools import lru_cache
from http import HTTPStatus
from typing import Any, Awaitable, Dict, Iterable, List, Optional, Tuple, Union
from typing import (
Any,
Awaitable,
Callable,
Dict,
Iterable,
List,
Optional,
Tuple,
Union,
)
import fastapi
import numpy as np
@@ -396,9 +406,20 @@ class InputFormat(Enum):
_MANAGER_OWNED_FIELDS = ("model_path", "served_model_name")
# Grace period from ShutdownReq to SIGKILL for each scheduler.
_SCHEDULER_EXIT_TIMEOUT_SECS = 15
class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
"""TokenizerManager is a process that tokenizes the text."""
# Set by whoever owns the event loop, and left None for Engine and grpc,
# which own no server. Class-level to leave the frozen __init__ alone.
_server_stop_hook: Optional[Callable[[], None]] = None
def set_server_stop_hook(self, hook: Callable[[], None]) -> None:
self._server_stop_hook = hook
@property
def serving_chat_class(self):
"""Return the serving chat class for OpenAI API.
@@ -3246,10 +3267,29 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# Ask schedulers to release resources in userspace and exit (see
# ShutdownReq), then wait for them before hard-killing the rest.
self._dispatch_to_scheduler(ShutdownReq())
deadline = time.monotonic() + 15
deadline = time.monotonic() + _SCHEDULER_EXIT_TIMEOUT_SECS
while time.monotonic() < deadline and collect_scheduler_processes():
time.sleep(0.1)
stragglers = [proc.pid for proc in collect_scheduler_processes()]
if stragglers:
# SIGKILL here lands mid-release,
# which is how GPU memory survives a shutdown. Name the pids.
logger.warning(
f"Schedulers still alive {_SCHEDULER_EXIT_TIMEOUT_SECS}s after "
f"ShutdownReq, killing them before they released: {stragglers}"
)
kill_process_tree(os.getpid(), include_parent=False, wait_timeout=60)
if self._server_stop_hook is not None:
# sys.exit() here raises SystemExit into the loop and kills it,
# so the ASGI server never runs its lifespan shutdown.
# The loop outlives this coroutine now, so drop our own tasks first;
# a pending handle_loop would be reported as destroyed-while-pending.
current = asyncio.current_task()
for task in self.asyncio_tasks:
if task is not current:
task.cancel()
self._server_stop_hook()
return
sys.exit(0)
def force_exit_handler(self):
+11 -1
View File
@@ -2007,6 +2007,7 @@ _GPU_IDLE_POLL_INTERVAL_SECS = 2.0
_GPU_IDLE_USED_MEMORY_THRESHOLD = 2 << 30 # 2 GiB
_GPU_RELEASE_TIMEOUT_SECS = 60.0
_GPU_RELEASE_POLL_INTERVAL_SECS = 0.5
_GPU_RELEASE_REPORT_THRESHOLD_SECS = 1.0
def _format_gib(num_bytes: Optional[int]) -> str:
@@ -2164,10 +2165,19 @@ def wait_for_gpu_release(
try:
gpu_indices = _visible_gpu_indices(pynvml)
pending = set(pids)
deadline = time.monotonic() + timeout
start = time.monotonic()
deadline = start + timeout
while True:
holders = _gpu_memory_holders(pynvml, gpu_indices, pending)
if not holders:
# Without this, a wait is indistinguishable from no wait.
waited = time.monotonic() - start
if waited >= _GPU_RELEASE_REPORT_THRESHOLD_SECS:
print(
f"[CI GPU Release] Waited {waited:.1f}s for"
f" {len(pending)} pid(s) to release.",
flush=True,
)
return
if time.monotonic() >= deadline:
print(