[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 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.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 can be set to 60
|
||||||
SGLANG_RINGBUFFER_WARNING_INTERVAL = int(
|
SGLANG_RINGBUFFER_WARNING_INTERVAL = int(
|
||||||
@@ -100,7 +101,9 @@ class ShmRingBuffer:
|
|||||||
# we are creating a buffer
|
# we are creating a buffer
|
||||||
self.is_creator = True
|
self.is_creator = True
|
||||||
self.shared_memory = shared_memory.SharedMemory(
|
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
|
# initialize the metadata section to 0
|
||||||
with memoryview(
|
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.custom_op import register_custom_op
|
||||||
from sglang.srt.utils.network import get_local_ip_auto
|
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_npu = is_npu()
|
||||||
_is_cpu = is_cpu()
|
_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):
|
with contextlib.suppress(OSError):
|
||||||
if rank == source_rank:
|
if rank == source_rank:
|
||||||
# create a shared memory segment
|
# 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
|
shm.buf[: len(magic_message)] = magic_message
|
||||||
torch.distributed.broadcast_object_list(
|
torch.distributed.broadcast_object_list(
|
||||||
[shm.name], src=ranks[source_rank], group=pg
|
[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.multimodal.evs import EVSEmbeddingResult
|
||||||
from sglang.srt.server_args import get_global_server_args
|
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 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
|
from sglang.utils import logger
|
||||||
|
|
||||||
_is_npu = is_npu()
|
_is_npu = is_npu()
|
||||||
@@ -1570,7 +1571,9 @@ class ShmPointerMMData:
|
|||||||
self.dtype = tensor.dtype
|
self.dtype = tensor.dtype
|
||||||
self.precomputed_hash = precomputed_hash
|
self.precomputed_hash = precomputed_hash
|
||||||
nbytes = tensor.numel() * tensor.element_size()
|
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:
|
try:
|
||||||
dst = torch.frombuffer(shm.buf, dtype=torch.uint8)
|
dst = torch.frombuffer(shm.buf, dtype=torch.uint8)
|
||||||
dst.copy_(tensor.view(torch.uint8).reshape(-1))
|
dst.copy_(tensor.view(torch.uint8).reshape(-1))
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ import torch
|
|||||||
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.server_args import get_global_server_args
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -67,7 +68,9 @@ def _pool_handle_cache_clear():
|
|||||||
|
|
||||||
class ShmSyncBuffer:
|
class ShmSyncBuffer:
|
||||||
def __init__(self, byte_size: int = 4):
|
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 = np.ndarray(1, dtype=np.float32, buffer=self.buffer.buf)
|
||||||
self.buffer_wrapper *= 0
|
self.buffer_wrapper *= 0
|
||||||
self.meta_data = {
|
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]}"
|
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() {
|
install_apt_packages() {
|
||||||
apt-get update || true
|
apt-get update || true
|
||||||
CI_APT_PACKAGES=(
|
CI_APT_PACKAGES=(
|
||||||
@@ -511,6 +522,7 @@ main() {
|
|||||||
configure_environment "$@"
|
configure_environment "$@"
|
||||||
detect_host
|
detect_host
|
||||||
kill_existing_processes
|
kill_existing_processes
|
||||||
|
cleanup_stale_shm
|
||||||
install_apt_packages
|
install_apt_packages
|
||||||
clean_site_packages
|
clean_site_packages
|
||||||
setup_pip_toolchain
|
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