[CI] Reclaim leaked /dev/shm segments on server startup (#28089)
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
7221be2cec
commit
cad43d3212
@@ -19,6 +19,7 @@ from zmq import IPV6 # type: ignore
|
||||
from zmq import SUB, SUBSCRIBE, XPUB, XPUB_VERBOSE, Context # type: ignore
|
||||
|
||||
from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto, get_open_port
|
||||
from sglang.srt.utils.stale_shm_cleanup import make_shm_name
|
||||
|
||||
# SGLANG_RINGBUFFER_WARNING_INTERVAL can be set to 60
|
||||
SGLANG_RINGBUFFER_WARNING_INTERVAL = int(
|
||||
@@ -100,7 +101,9 @@ class ShmRingBuffer:
|
||||
# we are creating a buffer
|
||||
self.is_creator = True
|
||||
self.shared_memory = shared_memory.SharedMemory(
|
||||
create=True, size=self.total_bytes_of_buffer
|
||||
create=True,
|
||||
size=self.total_bytes_of_buffer,
|
||||
name=make_shm_name("mq"),
|
||||
)
|
||||
# initialize the metadata section to 0
|
||||
with memoryview(
|
||||
|
||||
@@ -63,6 +63,7 @@ from sglang.srt.utils import (
|
||||
)
|
||||
from sglang.srt.utils.custom_op import register_custom_op
|
||||
from sglang.srt.utils.network import get_local_ip_auto
|
||||
from sglang.srt.utils.stale_shm_cleanup import make_shm_name
|
||||
|
||||
_is_npu = is_npu()
|
||||
_is_cpu = is_cpu()
|
||||
@@ -2444,7 +2445,9 @@ def in_the_same_node_as(pg: ProcessGroup, source_rank: int = 0) -> List[bool]:
|
||||
with contextlib.suppress(OSError):
|
||||
if rank == source_rank:
|
||||
# create a shared memory segment
|
||||
shm = shared_memory.SharedMemory(create=True, size=128)
|
||||
shm = shared_memory.SharedMemory(
|
||||
create=True, size=128, name=make_shm_name("nodecheck")
|
||||
)
|
||||
shm.buf[: len(magic_message)] = magic_message
|
||||
torch.distributed.broadcast_object_list(
|
||||
[shm.name], src=ranks[source_rank], group=pg
|
||||
|
||||
@@ -27,6 +27,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.multimodal.evs import EVSEmbeddingResult
|
||||
from sglang.srt.server_args import get_global_server_args
|
||||
from sglang.srt.utils import flatten_nested_list, is_npu, print_warning_once
|
||||
from sglang.srt.utils.stale_shm_cleanup import make_shm_name
|
||||
from sglang.utils import logger
|
||||
|
||||
_is_npu = is_npu()
|
||||
@@ -1570,7 +1571,9 @@ class ShmPointerMMData:
|
||||
self.dtype = tensor.dtype
|
||||
self.precomputed_hash = precomputed_hash
|
||||
nbytes = tensor.numel() * tensor.element_size()
|
||||
shm = shared_memory.SharedMemory(create=True, size=nbytes)
|
||||
shm = shared_memory.SharedMemory(
|
||||
create=True, size=nbytes, name=make_shm_name("mm")
|
||||
)
|
||||
try:
|
||||
dst = torch.frombuffer(shm.buf, dtype=torch.uint8)
|
||||
dst.copy_(tensor.view(torch.uint8).reshape(-1))
|
||||
|
||||
@@ -10,6 +10,7 @@ import torch
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.server_args import get_global_server_args
|
||||
from sglang.srt.utils.stale_shm_cleanup import make_shm_name
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -67,7 +68,9 @@ def _pool_handle_cache_clear():
|
||||
|
||||
class ShmSyncBuffer:
|
||||
def __init__(self, byte_size: int = 4):
|
||||
self.buffer = shared_memory.SharedMemory(create=True, size=byte_size)
|
||||
self.buffer = shared_memory.SharedMemory(
|
||||
create=True, size=byte_size, name=make_shm_name("sync")
|
||||
)
|
||||
self.buffer_wrapper = np.ndarray(1, dtype=np.float32, buffer=self.buffer.buf)
|
||||
self.buffer_wrapper *= 0
|
||||
self.meta_data = {
|
||||
|
||||
@@ -0,0 +1,123 @@
|
||||
"""Self-heal for leaked POSIX shared-memory segments in CI.
|
||||
|
||||
SGLang processes are torn down with SIGKILL (kill_process_tree, PDEATHSIG),
|
||||
which skips every Python-level unlink path, so /dev/shm segments accumulate
|
||||
until the tmpfs is full and the next scheduler init dies with SIGBUS.
|
||||
|
||||
Segments created through make_shm_name() embed the creator pid, which lets a
|
||||
later server startup safely unlink segments whose creator is gone. The sweep
|
||||
only runs in CI (single-tenant runner containers); on shared dev machines a
|
||||
pid check against another user's process is not authoritative, so we skip.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_SHM_DIR = Path("/dev/shm")
|
||||
_SGL_SHM_PREFIX = "sgl_shm"
|
||||
|
||||
|
||||
def make_shm_name(kind: str) -> str:
|
||||
"""Name a shared-memory segment so cleanup_stale_shm can identify and
|
||||
reclaim it after its creator process dies: sgl_shm_<kind>_<pid>_<rand>."""
|
||||
return f"{_SGL_SHM_PREFIX}_{kind}_{os.getpid()}_{uuid.uuid4().hex[:8]}"
|
||||
|
||||
|
||||
def _creator_pid(filename: str) -> int | None:
|
||||
pid = None
|
||||
if filename.startswith(f"{_SGL_SHM_PREFIX}_"):
|
||||
# sgl_shm_<kind>_<pid>_<rand>
|
||||
parts = filename.split("_")
|
||||
if len(parts) >= 4:
|
||||
try:
|
||||
pid = int(parts[-2])
|
||||
except ValueError:
|
||||
return None
|
||||
elif filename.startswith("multi_tokenizer_args_"):
|
||||
try:
|
||||
pid = int(filename.rsplit("_", 1)[-1])
|
||||
except ValueError:
|
||||
return None
|
||||
# os.kill(0, ...) / os.kill(-1, ...) probe process groups, not a process.
|
||||
if pid is not None and pid <= 0:
|
||||
return None
|
||||
return pid
|
||||
|
||||
|
||||
def _pid_alive(pid: int) -> bool:
|
||||
try:
|
||||
os.kill(pid, 0)
|
||||
return True
|
||||
except ProcessLookupError:
|
||||
return False
|
||||
except PermissionError:
|
||||
# Process exists but is owned by someone else.
|
||||
return True
|
||||
|
||||
|
||||
def cleanup_stale_shm() -> None:
|
||||
"""Unlink shared-memory segments whose creator process is dead.
|
||||
|
||||
CI-only: gated on SGLANG_IS_IN_CI because the pid-liveness check is only
|
||||
trustworthy when the container runs one job at a time. Best-effort: never
|
||||
raises, since a failed sweep must not block server startup.
|
||||
"""
|
||||
try:
|
||||
_cleanup_stale_shm_impl()
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"cleanup_stale_shm: sweep failed, continuing startup", exc_info=True
|
||||
)
|
||||
|
||||
|
||||
def _is_in_ci() -> bool:
|
||||
# Read the env var directly (same semantics as sglang.utils.is_in_ci) so
|
||||
# this module stays import-free and runnable by path from CI scripts
|
||||
# before sglang is installed.
|
||||
return os.environ.get("SGLANG_IS_IN_CI", "false").lower() in ("true", "1")
|
||||
|
||||
|
||||
def _cleanup_stale_shm_impl() -> None:
|
||||
if not _is_in_ci():
|
||||
return
|
||||
if not _SHM_DIR.is_dir():
|
||||
return
|
||||
|
||||
removed = 0
|
||||
freed_bytes = 0
|
||||
try:
|
||||
entries = list(_SHM_DIR.iterdir())
|
||||
except OSError as e:
|
||||
logger.warning("cleanup_stale_shm: cannot list %s, skipping: %s", _SHM_DIR, e)
|
||||
return
|
||||
for entry in entries:
|
||||
pid = _creator_pid(entry.name)
|
||||
if pid is None or pid == os.getpid() or _pid_alive(pid):
|
||||
# A recycled pid reads as alive, so pid-reuse degrades to
|
||||
# under-collection (segment leaks), never to deleting a live
|
||||
# segment. Keep that bias when changing this check.
|
||||
continue
|
||||
try:
|
||||
size = entry.stat().st_size
|
||||
entry.unlink()
|
||||
removed += 1
|
||||
freed_bytes += size
|
||||
except FileNotFoundError:
|
||||
pass # raced with another cleaner
|
||||
except OSError as e:
|
||||
logger.warning("cleanup_stale_shm: failed to remove %s: %s", entry.name, e)
|
||||
if removed:
|
||||
logger.info(
|
||||
"cleanup_stale_shm: removed %d stale segment(s), freed %.1f MiB",
|
||||
removed,
|
||||
freed_bytes / (1 << 20),
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
cleanup_stale_shm()
|
||||
@@ -122,6 +122,17 @@ kill_existing_processes() {
|
||||
mark_step_done "${FUNCNAME[0]}"
|
||||
}
|
||||
|
||||
cleanup_stale_shm() {
|
||||
# Reclaim /dev/shm segments leaked by SIGKILLed processes from earlier
|
||||
# jobs; leaked segments accumulate until the tmpfs fills and scheduler
|
||||
# init dies with SIGBUS. Runs right after killall so every dead creator's
|
||||
# segments are reclaimable. The module is dependency-free and runnable by
|
||||
# path, so this works before sglang is installed.
|
||||
SGLANG_IS_IN_CI=true python3 "${REPO_ROOT}/python/sglang/srt/utils/stale_shm_cleanup.py" || true
|
||||
|
||||
mark_step_done "${FUNCNAME[0]}"
|
||||
}
|
||||
|
||||
install_apt_packages() {
|
||||
apt-get update || true
|
||||
CI_APT_PACKAGES=(
|
||||
@@ -511,6 +522,7 @@ main() {
|
||||
configure_environment "$@"
|
||||
detect_host
|
||||
kill_existing_processes
|
||||
cleanup_stale_shm
|
||||
install_apt_packages
|
||||
clean_site_packages
|
||||
setup_pip_toolchain
|
||||
|
||||
@@ -0,0 +1,126 @@
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import unittest
|
||||
from multiprocessing import shared_memory
|
||||
from unittest.mock import patch
|
||||
|
||||
from sglang.srt.utils.stale_shm_cleanup import (
|
||||
_creator_pid,
|
||||
cleanup_stale_shm,
|
||||
make_shm_name,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
|
||||
register_cpu_ci(est_time=6, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
def _spawn_dead_pid() -> int:
|
||||
"""Return a pid that is guaranteed dead (already reaped)."""
|
||||
proc = subprocess.Popen([sys.executable, "-c", "pass"])
|
||||
proc.wait()
|
||||
return proc.pid
|
||||
|
||||
|
||||
class TestMakeShmName(unittest.TestCase):
|
||||
def test_embeds_pid_and_is_unique(self):
|
||||
a, b = make_shm_name("mm"), make_shm_name("mm")
|
||||
self.assertNotEqual(a, b)
|
||||
self.assertEqual(_creator_pid(a), os.getpid())
|
||||
|
||||
def test_creator_pid_parsing(self):
|
||||
self.assertEqual(_creator_pid("sgl_shm_mq_1234_abcd1234"), 1234)
|
||||
self.assertEqual(_creator_pid("multi_tokenizer_args_5678"), 5678)
|
||||
self.assertIsNone(_creator_pid("psm_deadbeef"))
|
||||
self.assertIsNone(_creator_pid("sgl_shm_garbage"))
|
||||
self.assertIsNone(_creator_pid("multi_tokenizer_args_notanint"))
|
||||
# Non-positive pids would make os.kill probe process groups.
|
||||
self.assertIsNone(_creator_pid("sgl_shm_mm_-1_abcd1234"))
|
||||
self.assertIsNone(_creator_pid("sgl_shm_mm_0_abcd1234"))
|
||||
|
||||
|
||||
@unittest.skipUnless(os.path.isdir("/dev/shm"), "requires /dev/shm")
|
||||
class TestCleanupStaleShm(unittest.TestCase):
|
||||
def _make_segment(self, name: str) -> str:
|
||||
shm = shared_memory.SharedMemory(create=True, size=4096, name=name)
|
||||
shm.close()
|
||||
self.addCleanup(self._unlink_quiet, name)
|
||||
return name
|
||||
|
||||
@staticmethod
|
||||
def _unlink_quiet(name: str):
|
||||
try:
|
||||
shared_memory.SharedMemory(name=name).unlink()
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
|
||||
def test_removes_dead_creator_keeps_live_and_foreign(self):
|
||||
dead_pid = _spawn_dead_pid()
|
||||
stale = self._make_segment(f"sgl_shm_mm_{dead_pid}_aaaa0000")
|
||||
live = self._make_segment(f"sgl_shm_mm_{os.getpid()}_bbbb0000")
|
||||
# Anonymous segments from other processes get psm_* names; the sweep
|
||||
# must never touch them even when their creator is dead.
|
||||
foreign = self._make_segment("psm_testforeign")
|
||||
|
||||
with patch.dict(os.environ, {"SGLANG_IS_IN_CI": "true"}):
|
||||
cleanup_stale_shm()
|
||||
|
||||
self.assertFalse(os.path.exists(f"/dev/shm/{stale}"))
|
||||
self.assertTrue(os.path.exists(f"/dev/shm/{live}"))
|
||||
self.assertTrue(os.path.exists(f"/dev/shm/{foreign}"))
|
||||
|
||||
def test_noop_outside_ci(self):
|
||||
dead_pid = _spawn_dead_pid()
|
||||
stale = self._make_segment(f"sgl_shm_mq_{dead_pid}_cccc0000")
|
||||
|
||||
with patch.dict(os.environ, {"SGLANG_IS_IN_CI": "false"}):
|
||||
cleanup_stale_shm()
|
||||
|
||||
self.assertTrue(os.path.exists(f"/dev/shm/{stale}"))
|
||||
|
||||
def test_shm_ring_buffer_uses_reclaimable_name(self):
|
||||
"""Bind the production call site: ShmRingBuffer must emit a
|
||||
pid-stamped name, or the leak this module fixes silently returns."""
|
||||
from sglang.srt.distributed.device_communicators.shm_broadcast import (
|
||||
ShmRingBuffer,
|
||||
)
|
||||
|
||||
buf = ShmRingBuffer(1, 64, 1)
|
||||
try:
|
||||
self.assertEqual(_creator_pid(buf.shared_memory.name), os.getpid())
|
||||
finally:
|
||||
buf.shared_memory.close()
|
||||
buf.shared_memory.unlink()
|
||||
|
||||
def test_run_by_path_without_sglang_importable(self):
|
||||
"""ci_install_dependency.sh runs the module by file path before
|
||||
sglang is installed; it must work with an empty PYTHONPATH."""
|
||||
import sglang.srt.utils.stale_shm_cleanup as mod
|
||||
|
||||
dead_pid = _spawn_dead_pid()
|
||||
stale = self._make_segment(f"sgl_shm_mm_{dead_pid}_eeee0000")
|
||||
|
||||
env = {k: v for k, v in os.environ.items() if k != "PYTHONPATH"}
|
||||
env["SGLANG_IS_IN_CI"] = "true"
|
||||
result = subprocess.run(
|
||||
[sys.executable, mod.__file__],
|
||||
env=env,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
cwd="/",
|
||||
)
|
||||
self.assertEqual(result.returncode, 0, result.stderr)
|
||||
self.assertFalse(os.path.exists(f"/dev/shm/{stale}"))
|
||||
|
||||
def test_multi_tokenizer_args_cleanup(self):
|
||||
dead_pid = _spawn_dead_pid()
|
||||
stale = self._make_segment(f"multi_tokenizer_args_{dead_pid}")
|
||||
|
||||
with patch.dict(os.environ, {"SGLANG_IS_IN_CI": "true"}):
|
||||
cleanup_stale_shm()
|
||||
|
||||
self.assertFalse(os.path.exists(f"/dev/shm/{stale}"))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user