[CI] Reclaim leaked /dev/shm segments on server startup (#28089)

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Kangyan-Zhou
2026-06-15 16:14:08 -07:00
committed by GitHub
co-authored by Claude Fable 5
parent 7221be2cec
commit cad43d3212
7 changed files with 277 additions and 4 deletions
@@ -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
+4 -1
View File
@@ -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()
+12
View File
@@ -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()