From 6580d5cd9a835f7fc66f88057e5bf6ef6506ea4a Mon Sep 17 00:00:00 2001 From: Tarang Khanna Date: Mon, 31 Aug 2026 15:44:45 +0700 Subject: [PATCH] weight cache: key daemon paths by GPU UUID (#36101) Co-authored-by: siyu Co-authored-by: Alex Nails --- python/sglang/srt/entrypoints/engine.py | 41 +++++++---- python/sglang/srt/environ.py | 10 +-- .../sglang/srt/model_executor/model_runner.py | 5 -- .../load_model_utils.py | 20 +----- python/sglang/srt/model_loader/loader.py | 13 +--- python/sglang/srt/server_args.py | 3 +- python/sglang/srt/weight_cache/daemon.py | 35 ++++++--- python/sglang/srt/weight_cache/ipc_loader.py | 12 +++- python/sglang/srt/weight_cache/protocol.py | 57 ++++++--------- .../model_loading/test_weight_cache_daemon.py | 43 ++++++----- .../test_weight_cache_protocol.py | 72 ++++++++++++++----- 11 files changed, 176 insertions(+), 135 deletions(-) diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index c2b5f4650..8a5167d07 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -101,6 +101,7 @@ from sglang.srt.observability.startup_time import build_engine_startup_time from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info from sglang.srt.parser.template_detection import resolve_auto_parsers from sglang.srt.parser.template_manager import TemplateManager +from sglang.srt.platforms import current_platform from sglang.srt.plugins import load_plugins from sglang.srt.runtime_context import ( get_disagg, @@ -138,6 +139,12 @@ from sglang.srt.utils.network import ( ) from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter from sglang.srt.utils.watchdog import SubprocessWatchdog +from sglang.srt.weight_cache.daemon import spawn_weight_cache_daemon +from sglang.srt.weight_cache.protocol import ( + cleanup_stale_daemon_files, + compute_local_gpu_id, + get_ready_path, +) from sglang.version import __version__ logger = logging.getLogger(__name__) @@ -715,19 +722,18 @@ class Engine(EngineScoreMixin, EngineBase): ) # Validate and clean up stale .ready/.sock files from prior runs. - # If a daemon is still alive at this rank, raise instead of clobbering. - from sglang.srt.weight_cache.daemon import spawn_weight_cache_daemon - from sglang.srt.weight_cache.protocol import ( - cleanup_stale_daemon_files, - compute_global_rank, - compute_local_gpu_id, - get_ready_path, - ) - + # If a daemon is still alive at this GPU, raise instead of clobbering. for pp_rank in pp_rank_range: for tp_rank in tp_rank_range: - global_rank = compute_global_rank(tp_size, pp_rank, tp_rank) - cleanup_stale_daemon_files(global_rank) + gpu_id = compute_local_gpu_id( + pp_rank, + tp_rank, + pp_size_per_node, + tp_size_per_node, + base_gpu_id=server_args.base_gpu_id, + gpu_id_step=server_args.gpu_id_step, + ) + cleanup_stale_daemon_files(current_platform.get_device_uuid(gpu_id)) for pp_rank in pp_rank_range: for tp_rank in tp_rank_range: @@ -759,8 +765,17 @@ class Engine(EngineScoreMixin, EngineBase): try: for pp_rank in pp_rank_range: for tp_rank in tp_rank_range: - global_rank = compute_global_rank(tp_size, pp_rank, tp_rank) - ready_path = get_ready_path(global_rank) + gpu_id = compute_local_gpu_id( + pp_rank, + tp_rank, + pp_size_per_node, + tp_size_per_node, + base_gpu_id=server_args.base_gpu_id, + gpu_id_step=server_args.gpu_id_step, + ) + ready_path = get_ready_path( + current_platform.get_device_uuid(gpu_id) + ) while not os.path.exists(ready_path): time.sleep(check_interval) if time.time() - start_time > timeout: diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 9231e8fd6..f83f2c637 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -1612,14 +1612,14 @@ class Envs: # Weight Cache Daemon # =================================================================== # Paths the daemon and the engine ranks it serves must agree on. Both are - # format templates and must keep the {global_rank} placeholder: each rank - # talks to the daemon on its own GPU, so a rank-independent path would point - # every rank at one daemon and map another rank's shard. + # format templates and must keep the {device_uuid} placeholder: each daemon + # is keyed by the physical GPU it runs on, so a GPU-independent path would + # let one job's client discover another job's daemon. SGLANG_WEIGHT_CACHE_SOCKET_TEMPLATE = EnvStr( - "/tmp/sglang_weight_cache_rank{global_rank}.sock" + "/tmp/sglang_weight_cache_{device_uuid}.sock" ) SGLANG_WEIGHT_CACHE_READY_TEMPLATE = EnvStr( - "/tmp/sglang_weight_cache_rank{global_rank}.ready" + "/tmp/sglang_weight_cache_{device_uuid}.ready" ) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 064ff0daf..a02221cbb 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -1145,13 +1145,8 @@ class ModelRunner: weight_cache_socket=get_model().weight_cache_socket, ) - # If the weight cache is enabled, override the load format to IPC_CACHE - # and derive the per-rank daemon socket. Idempotent across reloads. maybe_enable_ipc_weight_cache( load_config=self.load_config, - tp_size=self.ps.tp_size, - pp_rank=self.ps.pp_rank, - tp_rank=self.ps.tp_rank, ) if self.device == "cpu": self.model_config = adjust_config_with_unaligned_cpu_tp( diff --git a/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py b/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py index 6f22381bf..0ea3c517c 100644 --- a/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py +++ b/python/sglang/srt/model_executor/model_runner_components/load_model_utils.py @@ -236,16 +236,13 @@ def build_load_config( def maybe_enable_ipc_weight_cache( *, load_config: LoadConfig, - tp_size: int, - pp_rank: int, - tp_rank: int, ) -> None: """Switch ``load_config`` onto the IPC weight-cache path, in place. Overrides the load format to ``IPC_CACHE`` (remembering the original as the - disk fallback) and derives the per-rank daemon socket if unset. Idempotent: - the format swap is guarded on ``!= IPC_CACHE`` so a second call (e.g. a - weight reload) can't overwrite the captured fallback format. + disk fallback). Idempotent: the format swap is guarded on ``!= IPC_CACHE`` + so a second call (e.g. a weight reload) can't overwrite the captured + fallback format. """ if get_model().weight_cache_mode == "off": return @@ -254,17 +251,6 @@ def maybe_enable_ipc_weight_cache( load_config.fallback_load_format = load_config.load_format load_config.load_format = LoadFormat.IPC_CACHE - # Compute socket path using global rank (tp_size * pp_rank + tp_rank) so - # each daemon has a unique socket even across PP stages and nodes. - if load_config.weight_cache_socket is None: - from sglang.srt.weight_cache.protocol import ( - compute_global_rank, - get_socket_path, - ) - - global_rank = compute_global_rank(tp_size, pp_rank, tp_rank) - load_config.weight_cache_socket = get_socket_path(global_rank=global_rank) - def load_model_with_memory_saver( *, diff --git a/python/sglang/srt/model_loader/loader.py b/python/sglang/srt/model_loader/loader.py index 9b784d6d5..3211f0c41 100644 --- a/python/sglang/srt/model_loader/loader.py +++ b/python/sglang/srt/model_loader/loader.py @@ -4379,21 +4379,10 @@ def get_model_loader( if load_config.load_format == LoadFormat.IPC_CACHE: from sglang.srt.weight_cache.ipc_loader import IpcModelLoader - from sglang.srt.weight_cache.protocol import ( - compute_global_rank, - get_socket_path, - ) - if load_config.weight_cache_socket: - socket_path = load_config.weight_cache_socket - else: - - ps = get_parallel() - global_rank = compute_global_rank(ps.tp_size, ps.pp_rank, ps.tp_rank) - socket_path = get_socket_path(global_rank=global_rank) return IpcModelLoader( load_config=load_config, - socket_path=socket_path, + socket_path=load_config.weight_cache_socket, weight_cache_mode=load_config.weight_cache_mode, fallback_load_format=load_config.fallback_load_format, ) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 263bf3d35..24134256b 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -3619,7 +3619,8 @@ class ServerArgs: Optional[str], Arg( help="Unix socket path for weight cache daemon (client mode)." - "If not set, uses /tmp/sglang_weight_cache_rank{global_rank}.sock", + "If not set, derives the path from SGLANG_WEIGHT_CACHE_SOCKET_TEMPLATE " + "using the caller's physical GPU UUID.", ), NS("model"), ] = None diff --git a/python/sglang/srt/weight_cache/daemon.py b/python/sglang/srt/weight_cache/daemon.py index 8bb80ae33..077236cbe 100644 --- a/python/sglang/srt/weight_cache/daemon.py +++ b/python/sglang/srt/weight_cache/daemon.py @@ -173,12 +173,9 @@ class WeightCacheDaemon: self.revision = cfg.revision self.dist_init_method = dist_init_method - self.socket_path = get_socket_path( - compute_global_rank(self.tp_size, pp_rank, tp_rank) - ) - self.ready_path = get_ready_path( - compute_global_rank(self.tp_size, pp_rank, tp_rank) - ) + device_uuid = current_platform.get_device_uuid(gpu_id) + self.socket_path = get_socket_path(device_uuid) + self.ready_path = get_ready_path(device_uuid) self.model = None self.config: Optional[CacheConfig] = None @@ -735,8 +732,17 @@ def launch_weight_cache_daemons( # Validate and clean up stale .ready/.sock files from prior runs. for pp_rank in pp_rank_range: for tp_rank in tp_rank_range: - global_rank = compute_global_rank(cfg.tp_size, pp_rank, tp_rank) - cleanup_stale_daemon_files(global_rank, force=force) + gpu_id = compute_local_gpu_id( + pp_rank, + tp_rank, + pp_size_per_node, + tp_size_per_node, + base_gpu_id=cfg.base_gpu_id, + gpu_id_step=cfg.gpu_id_step, + ) + cleanup_stale_daemon_files( + current_platform.get_device_uuid(gpu_id), force=force + ) procs = [] for pp_rank in pp_rank_range: @@ -768,8 +774,15 @@ def launch_weight_cache_daemons( start_time = time.time() for pp_rank in pp_rank_range: for tp_rank in tp_rank_range: - global_rank = compute_global_rank(cfg.tp_size, pp_rank, tp_rank) - ready_path = get_ready_path(global_rank) + gpu_id = compute_local_gpu_id( + pp_rank, + tp_rank, + pp_size_per_node, + tp_size_per_node, + base_gpu_id=cfg.base_gpu_id, + gpu_id_step=cfg.gpu_id_step, + ) + ready_path = get_ready_path(current_platform.get_device_uuid(gpu_id)) while not os.path.exists(ready_path): time.sleep(check_interval) if time.time() - start_time > timeout: @@ -863,7 +876,7 @@ if __name__ == "__main__": else daemon_args.gpu_id ) cleanup_stale_daemon_files( - compute_global_rank(server_args.tp_size, daemon_args.pp_rank, tp_rank), + current_platform.get_device_uuid(gpu_id), force=daemon_args.force, ) run_weight_cache_daemon( diff --git a/python/sglang/srt/weight_cache/ipc_loader.py b/python/sglang/srt/weight_cache/ipc_loader.py index b4de2b825..3ebbae067 100644 --- a/python/sglang/srt/weight_cache/ipc_loader.py +++ b/python/sglang/srt/weight_cache/ipc_loader.py @@ -22,6 +22,7 @@ from sglang.srt.model_loader.loader import ( BaseModelLoader, _initialize_model, ) +from sglang.srt.platforms import current_platform from sglang.srt.runtime_context import get_exec, get_parallel from .protocol import ( @@ -29,6 +30,7 @@ from .protocol import ( check_ipc_quant_support, compute_env_stamp, get_quant_method_name, + get_socket_path, hash_quant_config, recv_msg, send_msg, @@ -67,7 +69,7 @@ class IpcModelLoader(BaseModelLoader): def __init__( self, load_config: LoadConfig, - socket_path: str, + socket_path: Optional[str] = None, fallback_loader_cls=None, weight_cache_mode: str = "client", fallback_load_format: str = "auto", @@ -103,7 +105,7 @@ class IpcModelLoader(BaseModelLoader): check_ipc_quant_support(quant_method, engine_quant_config, where="client") # Try to fetch state from daemon - cache_data = self._fetch_from_cache(model_config) + cache_data = self._fetch_from_cache(model_config, device_config) if cache_data is None: if self.weight_cache_mode == "daemon": @@ -446,7 +448,7 @@ class IpcModelLoader(BaseModelLoader): return model - def _fetch_from_cache(self, model_config) -> Optional[dict]: + def _fetch_from_cache(self, model_config, device_config) -> Optional[dict]: """Connect to daemon, validate config, fetch IPC handles. Returns the daemon response dict on success, None if the daemon is @@ -455,6 +457,10 @@ class IpcModelLoader(BaseModelLoader): """ import socket as socket_mod + if self.socket_path is None: + device_uuid = current_platform.get_device_uuid(int(device_config.gpu_id)) + self.socket_path = get_socket_path(device_uuid) + # Only connect to a real socket node owned by us: reject a symlink, a # plain file, or another user's socket planted at this /tmp path. An # absent socket means no daemon -> fall back to disk (return None). diff --git a/python/sglang/srt/weight_cache/protocol.py b/python/sglang/srt/weight_cache/protocol.py index d2f043b4c..1ae09eb40 100644 --- a/python/sglang/srt/weight_cache/protocol.py +++ b/python/sglang/srt/weight_cache/protocol.py @@ -269,12 +269,7 @@ def compute_env_stamp() -> Dict[str, str]: def compute_global_rank(tp_size: int, pp_rank: int, tp_rank: int) -> int: - """Single source of truth for the daemon rank formula. - - global_rank = tp_size * pp_rank + tp_rank, so each daemon gets a unique - socket/ready path even across PP stages and nodes. Every call site (engine, - loader, model_runner, daemon) must go through this so the copies can't drift. - """ + """Global rank for ``init_distributed_environment`` (tp_size * pp_rank + tp_rank).""" return tp_size * pp_rank + tp_rank @@ -302,38 +297,32 @@ def compute_local_gpu_id( ) -def _format_daemon_path(env_field, global_rank: int) -> str: - """Fill in a daemon path template, rejecting one that drops the rank. +def _format_daemon_path(env_field, device_uuid: str) -> str: + """Fill in a daemon path template, rejecting one that drops the GPU identity. The template is user-overridable, and ``str.format`` silently ignores a - missing placeholder. Every rank would then derive the same path and map the - shard belonging to whichever daemon got there first, so refuse up front - rather than serve wrong weights. + missing placeholder. Every physical GPU would then derive the same path, + letting one job's client discover another job's daemon, so refuse up + front rather than serve wrong weights. """ template = env_field.get() - if "{global_rank}" not in template: + if "{device_uuid}" not in template: raise ValueError( - f"{env_field.name}={template!r} must contain '{{global_rank}}': each " - f"rank needs its own path, and a rank-independent one would point " - f"every rank at a single daemon." + f"{env_field.name}={template!r} must contain '{{device_uuid}}': " + f"each physical GPU needs its own path, and a GPU-independent one " + f"would point every caller at a single daemon." ) - return template.format(global_rank=global_rank) + return template.format(device_uuid=device_uuid) -def get_socket_path(global_rank: int) -> str: - """Get the Unix socket path for a weight cache daemon. - - global_rank = tp_size * pp_rank + tp_rank - """ - return _format_daemon_path(envs.SGLANG_WEIGHT_CACHE_SOCKET_TEMPLATE, global_rank) +def get_socket_path(device_uuid: str) -> str: + """Get the Unix socket path for a weight cache daemon's physical GPU.""" + return _format_daemon_path(envs.SGLANG_WEIGHT_CACHE_SOCKET_TEMPLATE, device_uuid) -def get_ready_path(global_rank: int) -> str: - """Get the ready-file path for a weight cache daemon. - - global_rank = tp_size * pp_rank + tp_rank - """ - return _format_daemon_path(envs.SGLANG_WEIGHT_CACHE_READY_TEMPLATE, global_rank) +def get_ready_path(device_uuid: str) -> str: + """Get the ready-file path for a weight cache daemon's physical GPU.""" + return _format_daemon_path(envs.SGLANG_WEIGHT_CACHE_READY_TEMPLATE, device_uuid) def _read_ready_pid(ready_path: str) -> Optional[int]: @@ -359,8 +348,8 @@ def _is_pid_alive(pid: int) -> bool: return True -def cleanup_stale_daemon_files(global_rank: int, *, force: bool = False) -> None: - """Validate and clean up .ready/.sock files for a daemon rank. +def cleanup_stale_daemon_files(device_uuid: str, *, force: bool = False) -> None: + """Validate and clean up .ready/.sock files for a daemon's physical GPU. If the .ready file exists and the recorded PID is still alive, the daemon is still running — raise RuntimeError so the caller doesn't clobber it, @@ -369,8 +358,8 @@ def cleanup_stale_daemon_files(global_rank: int, *, force: bool = False) -> None If the PID is dead (or unreadable), the files are stale leftovers from a crashed/killed daemon and are safe to remove. """ - ready_path = get_ready_path(global_rank) - socket_path = get_socket_path(global_rank) + ready_path = get_ready_path(device_uuid) + socket_path = get_socket_path(device_uuid) if not os.path.exists(ready_path) and not os.path.exists(socket_path): return @@ -380,14 +369,14 @@ def cleanup_stale_daemon_files(global_rank: int, *, force: bool = False) -> None if pid is not None and _is_pid_alive(pid): if not force: raise RuntimeError( - f"Weight cache daemon for rank {global_rank} is already running " + f"Weight cache daemon for GPU {device_uuid} is already running " f"(pid={pid}, ready={ready_path}). Stop the existing daemon before " f"launching a new one, or pass force=True (--force) to kill it and " f"take over." ) logger.warning( f"[weight_cache] force takeover: killing existing daemon pid={pid} " - f"for rank {global_rank} and reclaiming its socket/ready files." + f"for GPU {device_uuid} and reclaiming its socket/ready files." ) try: os.kill(pid, signal.SIGKILL) diff --git a/test/registered/model_loading/test_weight_cache_daemon.py b/test/registered/model_loading/test_weight_cache_daemon.py index 6be3a01fc..f7a90b2eb 100644 --- a/test/registered/model_loading/test_weight_cache_daemon.py +++ b/test/registered/model_loading/test_weight_cache_daemon.py @@ -7,7 +7,9 @@ import unittest import requests import torch +from sglang.srt.platforms import current_platform from sglang.srt.utils import kill_process_tree +from sglang.srt.weight_cache.protocol import get_ready_path, get_socket_path from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.test_utils import ( DEFAULT_TARGET_MODEL_EAGLE_DP_ATTN, @@ -44,6 +46,11 @@ PROMPTS = [ ] +def _gpu_uuids(tp_size: int) -> list: + # Single-node, default base_gpu_id/gpu_id_step: rank i runs on physical GPU i. + return [current_platform.get_device_uuid(i) for i in range(tp_size)] + + @unittest.skipIf( torch.cuda.device_count() < 2, "TP=2 weight cache daemon test requires >=2 GPUs (skipped on the 1-gpu runner)", @@ -60,11 +67,11 @@ class TestWeightCacheDaemonTP2(CustomTestCase): cls.model = cls.model_override or DEFAULT_MODEL cls.base_url = DEFAULT_URL_FOR_TEST cls.tp_size = 2 + cls.gpu_uuids = _gpu_uuids(cls.tp_size) # Clean up stale ready/socket files from previous runs - for rank in range(cls.tp_size): - for suffix in (".ready", ".sock"): - path = f"/tmp/sglang_weight_cache_rank{rank}{suffix}" + for device_uuid in cls.gpu_uuids: + for path in (get_ready_path(device_uuid), get_socket_path(device_uuid)): if os.path.exists(path): os.unlink(path) @@ -86,13 +93,14 @@ class TestWeightCacheDaemonTP2(CustomTestCase): # Step 2: Wait for all daemon ready files timeout = DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH start = time.time() - for rank in range(cls.tp_size): - ready_path = f"/tmp/sglang_weight_cache_rank{rank}.ready" + for device_uuid in cls.gpu_uuids: + ready_path = get_ready_path(device_uuid) while not os.path.exists(ready_path): if time.time() - start > timeout: kill_process_tree(cls.daemon_process.pid) raise TimeoutError( - f"Weight cache daemon rank {rank} not ready within {timeout}s" + f"Weight cache daemon for GPU {device_uuid} not ready " + f"within {timeout}s" ) if cls.daemon_process.poll() is not None: raise RuntimeError( @@ -136,9 +144,8 @@ class TestWeightCacheDaemonTP2(CustomTestCase): os.unlink(path) except OSError: pass - for rank in range(getattr(cls, "tp_size", 2)): - for suffix in (".ready", ".sock"): - path = f"/tmp/sglang_weight_cache_rank{rank}{suffix}" + for device_uuid in getattr(cls, "gpu_uuids", ()): + for path in (get_ready_path(device_uuid), get_socket_path(device_uuid)): if os.path.exists(path): try: os.unlink(path) @@ -220,11 +227,11 @@ class TestWeightCacheDaemonTP1Smoke(CustomTestCase): cls.model = DEFAULT_MODEL cls.base_url = DEFAULT_URL_FOR_TEST cls.tp_size = 1 + cls.gpu_uuids = _gpu_uuids(cls.tp_size) # Clean up stale ready/socket files from previous runs. - for rank in range(cls.tp_size): - for suffix in (".ready", ".sock"): - path = f"/tmp/sglang_weight_cache_rank{rank}{suffix}" + for device_uuid in cls.gpu_uuids: + for path in (get_ready_path(device_uuid), get_socket_path(device_uuid)): if os.path.exists(path): os.unlink(path) @@ -245,13 +252,14 @@ class TestWeightCacheDaemonTP1Smoke(CustomTestCase): # Step 2: Wait for the daemon ready file. timeout = DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH start = time.time() - for rank in range(cls.tp_size): - ready_path = f"/tmp/sglang_weight_cache_rank{rank}.ready" + for device_uuid in cls.gpu_uuids: + ready_path = get_ready_path(device_uuid) while not os.path.exists(ready_path): if time.time() - start > timeout: kill_process_tree(cls.daemon_process.pid) raise TimeoutError( - f"Weight cache daemon rank {rank} not ready within {timeout}s" + f"Weight cache daemon for GPU {device_uuid} not ready " + f"within {timeout}s" ) if cls.daemon_process.poll() is not None: raise RuntimeError( @@ -294,9 +302,8 @@ class TestWeightCacheDaemonTP1Smoke(CustomTestCase): os.unlink(path) except OSError: pass - for rank in range(getattr(cls, "tp_size", 1)): - for suffix in (".ready", ".sock"): - path = f"/tmp/sglang_weight_cache_rank{rank}{suffix}" + for device_uuid in getattr(cls, "gpu_uuids", ()): + for path in (get_ready_path(device_uuid), get_socket_path(device_uuid)): if os.path.exists(path): try: os.unlink(path) diff --git a/test/registered/unit/model_loader/test_weight_cache_protocol.py b/test/registered/unit/model_loader/test_weight_cache_protocol.py index e1aa82da0..deb75ffd0 100644 --- a/test/registered/unit/model_loader/test_weight_cache_protocol.py +++ b/test/registered/unit/model_loader/test_weight_cache_protocol.py @@ -5,7 +5,7 @@ These cover the pure-Python logic that the GPU end-to-end test (test_weight_cache_daemon.py) cannot exercise cheaply: - length-prefixed socket framing (send_msg/recv_msg) over socketpair() - - CacheConfig fingerprint matching / (de)serialization + - CacheConfig compatibility matching / (de)serialization - quant-config hashing and method-name extraction - daemon spawn configuration and socket/ready path derivation - the IPC quantization allowlist (the gate that keeps silently-wrong @@ -25,6 +25,7 @@ from types import SimpleNamespace import torch +from sglang.srt.environ import envs from sglang.srt.weight_cache.protocol import ( IPC_QUANT_ALLOWLIST, CacheConfig, @@ -240,10 +241,25 @@ class TestGlobalRankAndPaths(CustomTestCase): self.assertEqual(compute_global_rank(tp_size=4, pp_rank=1, tp_rank=0), 4) self.assertEqual(compute_global_rank(tp_size=4, pp_rank=2, tp_rank=1), 9) - def test_socket_and_ready_paths_are_unique_per_rank(self): - self.assertNotEqual(get_socket_path(0), get_socket_path(1)) - self.assertTrue(get_socket_path(3).endswith("rank3.sock")) - self.assertTrue(get_ready_path(3).endswith("rank3.ready")) + def test_socket_and_ready_paths_are_unique_per_device_uuid(self): + self.assertNotEqual(get_socket_path("gpu-aaa"), get_socket_path("gpu-bbb")) + self.assertTrue(get_socket_path("gpu-aaa").endswith("gpu-aaa.sock")) + self.assertTrue(get_ready_path("gpu-aaa").endswith("gpu-aaa.ready")) + + def test_custom_template_with_device_uuid_placeholder_is_honored(self): + with envs.SGLANG_WEIGHT_CACHE_SOCKET_TEMPLATE.override( + "/custom/dir/{device_uuid}.custom-sock" + ): + self.assertEqual( + get_socket_path("gpu-aaa"), "/custom/dir/gpu-aaa.custom-sock" + ) + + def test_template_missing_device_uuid_placeholder_raises(self): + with envs.SGLANG_WEIGHT_CACHE_SOCKET_TEMPLATE.override( + "/tmp/sglang_weight_cache_rank{global_rank}.sock" + ): + with self.assertRaises(ValueError): + get_socket_path("gpu-aaa") def test_compute_local_gpu_id_honors_base_and_step(self): # Single-node TP=4: identity mapping rank -> gpu. @@ -368,12 +384,12 @@ class TestIpcQuantAllowlist(CustomTestCase): class TestCleanupStaleDaemonFiles(CustomTestCase): - # Use a rank far outside any realistic tp*pp layout so we never collide - # with a daemon that might actually be running on the test host. - RANK = 987654 + # A key no real daemon would ever compute, so this never collides with + # one that might actually be running on the test host. + KEY = "test-cleanup-stale-daemon-files" def _paths(self): - return get_ready_path(self.RANK), get_socket_path(self.RANK) + return get_ready_path(self.KEY), get_socket_path(self.KEY) def tearDown(self): for path in self._paths(): @@ -382,7 +398,7 @@ class TestCleanupStaleDaemonFiles(CustomTestCase): def test_no_files_is_noop(self): # Neither file present: must return quietly, not raise. - cleanup_stale_daemon_files(self.RANK) + cleanup_stale_daemon_files(self.KEY) def test_stale_files_without_live_pid_are_removed(self): ready_path, socket_path = self._paths() @@ -392,7 +408,7 @@ class TestCleanupStaleDaemonFiles(CustomTestCase): f.write("stale contents, no pid line\n") open(socket_path, "w").close() - cleanup_stale_daemon_files(self.RANK) + cleanup_stale_daemon_files(self.KEY) self.assertFalse(os.path.exists(ready_path)) self.assertFalse(os.path.exists(socket_path)) @@ -405,7 +421,7 @@ class TestCleanupStaleDaemonFiles(CustomTestCase): open(socket_path, "w").close() with self.assertRaises(RuntimeError): - cleanup_stale_daemon_files(self.RANK) + cleanup_stale_daemon_files(self.KEY) self.assertTrue(os.path.exists(ready_path)) self.assertTrue(os.path.exists(socket_path)) @@ -423,11 +439,11 @@ class TestCleanupStaleDaemonFiles(CustomTestCase): f.write(f"pid={child.pid}\n") open(socket_path, "w").close() - cleanup_stale_daemon_files(self.RANK, force=True) + cleanup_stale_daemon_files(self.KEY, force=True) self.assertFalse(os.path.exists(ready_path)) self.assertFalse(os.path.exists(socket_path)) - # The daemon holding the rank must have been killed. + # The daemon holding the identity must have been killed. self.assertEqual(child.wait(timeout=5), -9) finally: if child.poll() is None: @@ -442,7 +458,7 @@ class TestDaemonModeRefusesDiskLoad(CustomTestCase): pointing the loader at a socket path that does not exist. """ - RANK = 987655 + KEY = "test-daemon-mode-refuses-disk-load" def _model_config(self): from types import SimpleNamespace @@ -464,7 +480,7 @@ class TestDaemonModeRefusesDiskLoad(CustomTestCase): from sglang.srt.configs.load_config import LoadConfig, LoadFormat from sglang.srt.weight_cache.ipc_loader import IpcModelLoader - missing_socket = get_socket_path(self.RANK) + missing_socket = get_socket_path(self.KEY) if os.path.exists(missing_socket): os.unlink(missing_socket) @@ -481,6 +497,30 @@ class TestDaemonModeRefusesDiskLoad(CustomTestCase): # fall through to a disk load. self.assertIn("daemon", str(ctx.exception).lower()) + def test_default_discovery_queries_device_uuid_for_the_caller_gpu(self): + """Without an explicit --weight-cache-socket (the production default), + discovery must derive the socket from the caller's own gpu_id via + get_device_uuid -- not silently substitute GPU 0.""" + from unittest import mock + + from sglang.srt.configs.load_config import LoadConfig, LoadFormat + from sglang.srt.weight_cache.ipc_loader import IpcModelLoader + + loader = IpcModelLoader( + load_config=LoadConfig(load_format=LoadFormat.IPC_CACHE), + weight_cache_mode="client", + fallback_load_format="auto", + ) + with mock.patch( + "sglang.srt.platforms.current_platform.get_device_uuid", + return_value="gpu-under-test", + ) as get_uuid: + result = loader._fetch_from_cache( + self._model_config(), SimpleNamespace(gpu_id=5) + ) + get_uuid.assert_called_once_with(5) + self.assertIsNone(result) # no real daemon at that socket -> absent + if __name__ == "__main__": unittest.main()