From f9c14e6bd4e17653ed2727c94010aa02ee84d794 Mon Sep 17 00:00:00 2001 From: Mike Qiu Date: Sat, 25 Jul 2026 14:31:21 +0800 Subject: [PATCH] [FEAT] Support fast engine recovery through weight cache (#27139) Signed-off-by: Michael Qiu Co-authored-by: liusy58 Co-authored-by: Alex Nails --- .github/CODEOWNERS | 1 + python/sglang/srt/configs/load_config.py | 6 + python/sglang/srt/entrypoints/engine.py | 236 ++++- python/sglang/srt/entrypoints/http_server.py | 1 + .../scheduler_components/weight_updater.py | 18 + .../sglang/srt/model_executor/model_runner.py | 13 + .../load_model_utils.py | 49 + .../model_runner_components/weight_updater.py | 19 + python/sglang/srt/model_loader/loader.py | 22 + python/sglang/srt/ray/http_server.py | 1 + python/sglang/srt/server_args.py | 59 ++ python/sglang/srt/weight_cache/__init__.py | 11 + python/sglang/srt/weight_cache/daemon.py | 944 ++++++++++++++++++ python/sglang/srt/weight_cache/ipc_loader.py | 565 +++++++++++ python/sglang/srt/weight_cache/protocol.py | 380 +++++++ test/manual/test_weight_cache_e2e.py | 377 +++++++ .../model_loading/test_weight_cache_daemon.py | 338 +++++++ .../test_weight_cache_protocol.py | 366 +++++++ 18 files changed, 3404 insertions(+), 2 deletions(-) create mode 100644 python/sglang/srt/weight_cache/__init__.py create mode 100644 python/sglang/srt/weight_cache/daemon.py create mode 100644 python/sglang/srt/weight_cache/ipc_loader.py create mode 100644 python/sglang/srt/weight_cache/protocol.py create mode 100644 test/manual/test_weight_cache_e2e.py create mode 100644 test/registered/model_loading/test_weight_cache_daemon.py create mode 100644 test/registered/unit/model_loader/test_weight_cache_protocol.py diff --git a/.github/CODEOWNERS b/.github/CODEOWNERS index d930516af..7aa2f9d21 100644 --- a/.github/CODEOWNERS +++ b/.github/CODEOWNERS @@ -65,6 +65,7 @@ /python/sglang/srt/ray @Qiaolin-Yu @xyuzh /python/sglang/srt/speculative @Ying1123 @merrymercy @hnyls2002 @Qiaolin-Yu /python/sglang/srt/utils/hf_transformers @JustinTong0323 +/python/sglang/srt/weight_cache @liusy58 @QiuMike /sgl-kernel @ispobock @BBuf @yizhang2077 @merrymercy @FlamingoPg @HaiShaw /sgl-kernel/csrc/musa @yeahdongcn /sgl-model-gateway @slin1237 @CatherineSue diff --git a/python/sglang/srt/configs/load_config.py b/python/sglang/srt/configs/load_config.py index 4a27652b9..c70ca6bc8 100644 --- a/python/sglang/srt/configs/load_config.py +++ b/python/sglang/srt/configs/load_config.py @@ -34,6 +34,7 @@ class LoadFormat(str, enum.Enum): FASTSAFETENSORS = "fastsafetensors" PRIVATE = "private" RUNAI_STREAMER = "runai_streamer" + IPC_CACHE = "ipc_cache" @dataclass @@ -104,6 +105,11 @@ class LoadConfig: # For multi-layer MTP draft_model_idx: Optional[int] = None + # Weight cache daemon options + weight_cache_mode: str = "off" # "off", "daemon", "client" + weight_cache_socket: Optional[str] = None # Path to daemon socket (for client mode) + fallback_load_format: Union[str, "LoadFormat"] = LoadFormat.AUTO + def __post_init__(self): model_loader_extra_config = self.model_loader_extra_config or {} if isinstance(model_loader_extra_config, str): diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index 888542300..25a80696c 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -27,6 +27,8 @@ import multiprocessing as mp import os import random import signal +import subprocess +import sys import tempfile import threading import time @@ -110,7 +112,12 @@ from sglang.srt.utils import ( set_ulimit, ) from sglang.srt.utils.msgspec_utils import msgspec_to_builtins -from sglang.srt.utils.network import get_zmq_socket, is_port_available +from sglang.srt.utils.network import ( + NetworkAddress, + get_free_port, + get_zmq_socket, + is_port_available, +) from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter from sglang.srt.utils.watchdog import SubprocessWatchdog from sglang.version import __version__ @@ -238,6 +245,7 @@ class Engine(EngineScoreMixin, EngineBase): port_args, scheduler_init_result, subprocess_watchdog, + weight_cache_daemon_procs, ) = self._launch_subprocesses( server_args=server_args, init_tokenizer_manager_func=self.init_tokenizer_manager_func, @@ -247,6 +255,10 @@ class Engine(EngineScoreMixin, EngineBase): self.tokenizer_manager = tokenizer_manager self.template_manager = template_manager self._scheduler_init_result = scheduler_init_result + # Engine-spawned weight cache daemons owned by *this* instance (empty + # unless --weight-cache-mode daemon). Kept per-instance so two Engines + # in one process each reap only their own daemons in shutdown(). + self._weight_cache_daemon_procs = weight_cache_daemon_procs if tokenizer_manager is not None: tokenizer_manager._subprocess_watchdog = subprocess_watchdog self.port_args = port_args @@ -588,6 +600,208 @@ class Engine(EngineScoreMixin, EngineBase): ret = self.loop.run_until_complete(generator.__anext__()) return ret + @classmethod + def _launch_weight_cache_daemons(cls, server_args: ServerArgs): + """Launch weight cache daemon processes for this node's PP×TP ranks. + + All daemon processes join the same NCCL distributed group so that + TP-sharded model loading works correctly. Each daemon holds its + rank's weight shard in GPU memory and serves IPC handles. + + Lifecycle: these daemons are *co-terminal* with the engine. They are + children of this process (kill_itself_when_parent_died installs + PR_SET_PDEATHSIG) and are gracefully reaped in ``shutdown()``. They do + NOT persist across engine restarts, so ``--weight-cache-mode daemon`` + on its own does not deliver a faster restart -- the first start is in + fact slower (disk-load into the daemon plus the IPC handshake). The + fast-recovery story is the standalone launcher + (``python -m sglang.srt.weight_cache.daemon``) plus + ``--weight-cache-mode client``, where the daemon outlives the engine. + """ + if server_args.dp_size > 1: + raise ValueError( + "Weight cache daemon mode does not support dp_size > 1. " + "Please set --dp-size 1 when using --weight-cache-mode daemon." + ) + + # Multi-node needs an explicit rendezvous address; otherwise each node + # picks its own local 127.0.0.1 port (below) and the per-node daemons + # can never form the joint process group. + if server_args.nnodes > 1 and not server_args.dist_init_addr: + raise ValueError( + "Multi-node weight cache daemons (nnodes > 1) require " + "--dist-init-addr so all nodes rendezvous at the same endpoint." + ) + + tp_size = server_args.tp_size + + pp_rank_range, tp_rank_range, pp_size_per_node, tp_size_per_node = ( + _calculate_rank_ranges( + server_args.nnodes, + server_args.pp_size, + tp_size, + server_args.node_rank, + ) + ) + + # Build the distributed init method (multi-node uses the user-provided + # dist_init_addr so all nodes reach the same endpoint). + if server_args.dist_init_addr: + host, port = server_args.dist_init_addr.rsplit(":", 1) + dist_init_method = f"tcp://{host}:{port}" + else: + # Fresh free port for the daemons' own rendezvous, not the engine's + # nccl_port: a pinned --nccl-port would otherwise collide with the + # engine's own NCCL TCPStore. + dist_init_method = NetworkAddress("127.0.0.1", get_free_port()).to_tcp() + + num_daemons = len(pp_rank_range) * len(tp_rank_range) + daemon_procs = [] + logger.info( + f"Launching {num_daemons} weight cache daemon(s) on node " + f"{server_args.node_rank} for model={server_args.model_path}, " + f"pp_ranks={pp_rank_range.start}..{pp_rank_range.stop - 1}, " + f"tp_ranks={tp_rank_range.start}..{tp_rank_range.stop - 1}, " + f"dist_init_method={dist_init_method}" + ) + + # 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.protocol import ( + cleanup_stale_daemon_files, + compute_global_rank, + compute_local_gpu_id, + get_ready_path, + ) + + 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) + + for pp_rank in pp_rank_range: + for tp_rank in tp_rank_range: + 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, + ) + cmd = [ + sys.executable, + "-m", + "sglang.srt.weight_cache.daemon", + "--model-path", + server_args.model_path, + "--gpu-id", + str(gpu_id), + "--tp-size", + str(tp_size), + "--tp-rank", + str(tp_rank), + "--pp-size", + str(server_args.pp_size), + "--pp-rank", + str(pp_rank), + "--dp-size", + "1", + "--ep-size", + str(server_args.ep_size), + "--load-format", + server_args.load_format, + "--dtype", + server_args.dtype, + "--dist-init-method", + dist_init_method, + ] + if server_args.quantization: + cmd += ["--quantization", server_args.quantization] + if ( + server_args.model_loader_extra_config + and server_args.model_loader_extra_config != "{}" + ): + cmd += [ + "--model-loader-extra-config", + server_args.model_loader_extra_config, + ] + if server_args.trust_remote_code: + cmd += ["--trust-remote-code"] + if server_args.revision: + cmd += ["--revision", server_args.revision] + + proc = subprocess.Popen(cmd) + daemon_procs.append(proc) + + # Wait for all daemons to be ready (ready file exists). On any failure + # (readiness timeout or a daemon exiting early) terminate the siblings + # we already spawned before propagating, so a partial launch does not + # leak GPU-resident daemons. + timeout = server_args.weight_cache_timeout + check_interval = 2 + start_time = time.time() + 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) + while not os.path.exists(ready_path): + time.sleep(check_interval) + if time.time() - start_time > timeout: + raise TimeoutError( + f"Weight cache daemon for pp_rank={pp_rank} " + f"tp_rank={tp_rank} did not become ready " + f"within {timeout}s" + ) + # Check if daemon process is still alive + for p in daemon_procs: + if p.poll() is not None: + raise RuntimeError( + f"Weight cache daemon (pid={p.pid}) exited prematurely " + f"with code {p.returncode}" + ) + logger.info( + f"Weight cache daemon for pp_rank={pp_rank} " + f"tp_rank={tp_rank} is ready" + ) + except BaseException: + cls._terminate_weight_cache_daemons(daemon_procs) + raise + + logger.info( + f"All {num_daemons} weight cache daemons on node " + f"{server_args.node_rank} are ready" + ) + return daemon_procs + + @staticmethod + def _terminate_weight_cache_daemons(procs, timeout: float = 10.0): + """Gracefully stop engine-spawned weight cache daemons. + + Send SIGTERM first so each daemon's signal handler can unlink its + ``.sock``/``.ready`` files, then SIGKILL any straggler. This matters + because ``shutdown()`` otherwise reaps children via + ``kill_process_tree`` (SIGKILL), which would skip that cleanup and + leave stale files that make the next client-mode boot fail with a + confusing "socket exists but connection refused" instead of a clean + "no daemon" path. + """ + if not procs: + return + for p in procs: + if p.poll() is None: + p.terminate() # SIGTERM -> daemon cleanup handler runs + for p in procs: + try: + p.wait(timeout=timeout) + except subprocess.TimeoutExpired: + logger.warning( + f"Weight cache daemon (pid={p.pid}) did not exit within " + f"{timeout}s of SIGTERM; sending SIGKILL." + ) + p.kill() + @classmethod def _launch_scheduler_processes( cls, @@ -779,7 +993,7 @@ class Engine(EngineScoreMixin, EngineBase): """Launch the TokenizerManager in the main process, the Scheduler in a subprocess, and the DetokenizerManager in another subprocess. Returns: - Tuple of (tokenizer_manager, template_manager, port_args, scheduler_init_result, subprocess_watchdog). + Tuple of (tokenizer_manager, template_manager, port_args, scheduler_init_result, subprocess_watchdog, weight_cache_daemon_procs). """ # Configure global environment configure_logger(server_args) @@ -820,6 +1034,13 @@ class Engine(EngineScoreMixin, EngineBase): ): resolve_auto_parsers(server_args) + # Launch daemons (daemon mode only). Handles are threaded back to the + # owning Engine instance (not a class attr) so two Engines in one process + # don't clobber each other's daemon list. + weight_cache_daemon_procs: List = [] + if server_args.weight_cache_mode == "daemon": + weight_cache_daemon_procs = cls._launch_weight_cache_daemons(server_args) + # Launch scheduler processes scheduler_init_result, scheduler_procs = cls._launch_scheduler_processes( server_args, port_args, run_scheduler_process_func @@ -846,6 +1067,7 @@ class Engine(EngineScoreMixin, EngineBase): port_args, scheduler_init_result, None, + weight_cache_daemon_procs, ) launch_dummy_health_check_server( @@ -859,6 +1081,7 @@ class Engine(EngineScoreMixin, EngineBase): port_args, scheduler_init_result, None, + weight_cache_daemon_procs, ) # Launch detokenizer process(es) — optionally fronted by a router when @@ -906,6 +1129,7 @@ class Engine(EngineScoreMixin, EngineBase): port_args, scheduler_init_result, subprocess_watchdog, + weight_cache_daemon_procs, ) def shutdown(self): @@ -923,6 +1147,14 @@ class Engine(EngineScoreMixin, EngineBase): send_to_rpc.close(linger=0) self.send_to_rpc = None + # Gracefully stop weight cache daemons *before* the blanket + # kill_process_tree below, so their SIGTERM handlers can unlink the + # .sock/.ready files instead of being SIGKILLed and leaving stale state. + daemon_procs = getattr(self, "_weight_cache_daemon_procs", None) + if daemon_procs: + self._terminate_weight_cache_daemons(daemon_procs) + self._weight_cache_daemon_procs = [] + kill_process_tree(os.getpid(), include_parent=False, wait_timeout=60) def __enter__(self): diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index 51c85e630..fd216d329 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -2675,6 +2675,7 @@ def launch_server( port_args, scheduler_init_result, subprocess_watchdog, + _weight_cache_daemon_procs, ) = Engine._launch_subprocesses( server_args=server_args, init_tokenizer_manager_func=init_tokenizer_manager_func, diff --git a/python/sglang/srt/managers/scheduler_components/weight_updater.py b/python/sglang/srt/managers/scheduler_components/weight_updater.py index 7013ce04f..653b28c41 100644 --- a/python/sglang/srt/managers/scheduler_components/weight_updater.py +++ b/python/sglang/srt/managers/scheduler_components/weight_updater.py @@ -183,6 +183,22 @@ class SchedulerWeightUpdaterManager: parameter = self.tp_worker.get_weights_by_name(recv_req) return GetWeightsByNameReqOutput(parameter=parameter) + def _assert_weight_cache_inactive(self, op: str) -> None: + """Reject freeing/restoring model weights while the CUDA IPC weight + cache is active: the weights are shared with the daemon via CUDA IPC, so + freeing them would leave the daemon and every peer pointing at released + memory. + """ + mode = self.tp_worker.model_runner.server_args.weight_cache_mode + if mode != "off": + raise RuntimeError( + f"[weight_cache] {op} of model weights is not supported while the " + f"weight cache is active (--weight-cache-mode {mode}): the weights " + f"are shared with the daemon via CUDA IPC, so freeing them would " + f"corrupt the daemon's master copy and every co-attached engine. " + f"Restart with --weight-cache-mode off to use this operation." + ) + def release_memory_occupation(self, recv_req: ReleaseMemoryOccupationReqInput): assert ( self.is_fully_idle() @@ -215,6 +231,7 @@ class SchedulerWeightUpdaterManager: self.flush_cache() if GPU_MEMORY_TYPE_WEIGHTS in tags: + self._assert_weight_cache_inactive("release_memory_occupation") self.stashed_model_static_state = _export_static_state( self.tp_worker.model_runner.model ) @@ -241,6 +258,7 @@ class SchedulerWeightUpdaterManager: self.memory_saver_adapter.resume(GPU_MEMORY_TYPE_CUDA_GRAPH) if GPU_MEMORY_TYPE_WEIGHTS in tags: + self._assert_weight_cache_inactive("resume_memory_occupation") self.memory_saver_adapter.resume(GPU_MEMORY_TYPE_WEIGHTS) torch.distributed.barrier(self.tp_cpu_group) _import_static_state( diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 207946c64..2d814f085 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -128,6 +128,7 @@ from sglang.srt.model_executor.model_runner_components.load_model_utils import ( load_kv_cache_scales, load_model_with_memory_saver, maybe_downgrade_dtype_for_legacy_gpu, + maybe_enable_ipc_weight_cache, maybe_register_debug_tensor_dump_hook, maybe_trigger_remote_instance_nccl_send_group, report_online_quantization, @@ -987,6 +988,18 @@ class ModelRunner: remote_instance_weight_transporter_engine=self.remote_instance_weight_transporter.engine, remote_instance_weight_transporter_session_id=self.remote_instance_weight_transporter.session_id, draft_model_idx=self.draft_model_idx, + weight_cache_mode=self.server_args.weight_cache_mode, + weight_cache_socket=self.server_args.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, + server_args=self.server_args, + 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 c55ba8007..855994aaf 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 @@ -172,6 +172,8 @@ def build_load_config( remote_instance_weight_transporter_engine: Any, remote_instance_weight_transporter_session_id: str, draft_model_idx: Optional[int], + weight_cache_mode: str, + weight_cache_socket: Optional[str], ) -> LoadConfig: from sglang.srt.configs.modelopt_config import ModelOptConfig @@ -199,9 +201,45 @@ def build_load_config( modelopt_config=modelopt_config, rl_quant_profile=server_args.rl_quant_profile, draft_model_idx=draft_model_idx, + weight_cache_mode=weight_cache_mode, + weight_cache_socket=weight_cache_socket, ) +def maybe_enable_ipc_weight_cache( + *, + load_config: LoadConfig, + server_args: ServerArgs, + 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. + """ + if server_args.weight_cache_mode == "off": + return + + if load_config.load_format != LoadFormat.IPC_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( *, server_args: ServerArgs, @@ -218,6 +256,17 @@ def load_model_with_memory_saver( enable_cpu_backup = server_args.enable_weights_cpu_backup or ( is_draft_worker and server_args.enable_draft_weights_cpu_backup ) + + # In zero-copy IPC mode, the weights are shared with the daemon via + # CUDA IPC and must not be offloaded/reloaded by the memory saver. + is_ipc_zero_copy = server_args.weight_cache_mode != "off" + if is_ipc_zero_copy and enable_cpu_backup: + logger.warning( + "[ModelRunner] Disabling weights CPU backup in zero-copy IPC mode — " + "IPC-mapped weights cannot be offloaded to CPU." + ) + enable_cpu_backup = False + remote_instance_weight_info = None with memory_saver_adapter.region( GPU_MEMORY_TYPE_WEIGHTS, diff --git a/python/sglang/srt/model_executor/model_runner_components/weight_updater.py b/python/sglang/srt/model_executor/model_runner_components/weight_updater.py index 1cfa8441b..ecbfa7a71 100644 --- a/python/sglang/srt/model_executor/model_runner_components/weight_updater.py +++ b/python/sglang/srt/model_executor/model_runner_components/weight_updater.py @@ -123,6 +123,21 @@ class WeightUpdater: logger.error(message) return False, message + def _assert_weight_cache_inactive(self: WeightUpdater, op: str) -> None: + """Reject weight mutations while the CUDA IPC weight cache is active: + param.data is the daemon's master copy shared with every co-attached + engine, so an in-place update would silently corrupt them all. + """ + mode = self.get_model_runner().server_args.weight_cache_mode + if mode != "off": + raise RuntimeError( + f"[weight_cache] {op} is not supported while the weight cache is " + f"active (--weight-cache-mode {mode}): model weights are shared " + f"with the daemon via CUDA IPC, so mutating them in place would " + f"corrupt the daemon's master copy and every co-attached engine. " + f"Restart with --weight-cache-mode off to use this operation." + ) + def update_weights_from_disk( self: WeightUpdater, model_path: str, @@ -131,6 +146,7 @@ class WeightUpdater: recapture_cuda_graph: bool = False, ) -> tuple[bool, str]: """Update engine weights in-place from the disk.""" + self._assert_weight_cache_inactive("update_weights_from_disk") error = _unsupported_derived_weight_cache_error() if error is not None: return False, error @@ -220,6 +236,7 @@ class WeightUpdater: dtype: the data type of the parameter to be updated. shape: the shape of the parameter to be updated. """ + self._assert_weight_cache_inactive("update_weights_from_distributed") error = _unsupported_derived_weight_cache_error() if error is not None: return False, error @@ -309,6 +326,7 @@ class WeightUpdater: return False, error monkey_patch_torch_reductions() + self._assert_weight_cache_inactive("update_weights_from_tensor") if load_format == "flattened_bucket": # Handle flattened bucket format return self._update_weights_from_flattened_bucket( @@ -368,6 +386,7 @@ class WeightUpdater: def update_weights_from_ipc(self: WeightUpdater, recv_req): """Update weights from IPC for checkpoint-engine integration.""" + self._assert_weight_cache_inactive("update_weights_from_ipc") error = _unsupported_derived_weight_cache_error() if error is not None: return False, error diff --git a/python/sglang/srt/model_loader/loader.py b/python/sglang/srt/model_loader/loader.py index 13c0c07d4..b673cb670 100644 --- a/python/sglang/srt/model_loader/loader.py +++ b/python/sglang/srt/model_loader/loader.py @@ -3319,4 +3319,26 @@ def get_model_loader( if load_config.load_format == LoadFormat.RUNAI_STREAMER: return RunaiModelStreamerLoader(load_config) + 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: + from sglang.srt.runtime_context import get_parallel + + 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, + weight_cache_mode=load_config.weight_cache_mode, + fallback_load_format=load_config.fallback_load_format, + ) + return DefaultModelLoader(load_config) diff --git a/python/sglang/srt/ray/http_server.py b/python/sglang/srt/ray/http_server.py index 7d7e23567..15d56cb6f 100644 --- a/python/sglang/srt/ray/http_server.py +++ b/python/sglang/srt/ray/http_server.py @@ -52,6 +52,7 @@ def launch_server( port_args, scheduler_init_result, subprocess_watchdog, + _weight_cache_daemon_procs, ) = RayEngine._launch_subprocesses( server_args, init_tokenizer_manager_func=init_tokenizer_manager_func, diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index a4ae73494..8db35a082 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -124,6 +124,10 @@ LOAD_FORMAT_CHOICES = [ "private", "runai_streamer", ] +# NOTE: LoadFormat.IPC_CACHE intentionally has no public --load-format choice. +# It is an internal dispatch format set automatically by ModelRunner when the +# weight cache is enabled (weight_cache_mode != "off"). Exposing it as a CLI +# choice let users create contradictory combos (see _handle_load_format). # TODO: this list should likely contain only methods that support online quantization, or that support using custom quantization classes compatible with a given `quant_method` in config.json. # Some of the choices here do NOT support online quantization. @@ -3302,6 +3306,36 @@ class ServerArgs: NS("exec.features"), ] = False + weight_cache_mode: A[ + str, + Arg( + help="Weight cache mode. 'off': normal disk loading. " + "'daemon': launch weight cache daemon (holds weights in GPU memory). " + "Engine-spawned daemons are co-terminal with the engine and do NOT " + "persist across restarts, so this alone does not speed up restart " + "(the first start is slower). For fast recovery, run the standalone " + "daemon (python -m sglang.srt.weight_cache.daemon) and connect with " + "'client'. 'client': connect to existing daemon and load via IPC.", + choices=["off", "daemon", "client"], + ), + NS("model"), + ] = "off" + weight_cache_socket: A[ + 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", + ), + NS("model"), + ] = None + weight_cache_timeout: A[ + int, + Arg( + help="Timeout in seconds for weight cache daemon readiness (default: 1800).", + ), + NS("model"), + ] = 1800 + # ------------------------------------------------------------------------- # Custom hooks, probe, and plugins # ------------------------------------------------------------------------- @@ -6926,6 +6960,31 @@ class ServerArgs: self.validate_transfer_engine() ) + # "ipc_cache" is an internal-only load format: ModelRunner sets it + # automatically when the weight cache is enabled, and it is not a public + # --load-format choice. Setting it directly is always wrong (no daemon is + # launched, and fallback_load_format inherits a nonsensical format), so + # reject it and point at the knob (defense-in-depth; the CLI already + # rejects it via LOAD_FORMAT_CHOICES). + if self.load_format == "ipc_cache": + raise ValueError( + "load_format='ipc_cache' is an internal-only format and must not " + "be set directly. Enable the weight cache via --weight-cache-mode " + "client (connect to an existing daemon) or daemon (launch one); " + "that selects IPC loading automatically." + ) + + # Speculative decoding loads an extra draft model whose weights the + # daemon does not export, so refuse the combination up front instead of + # failing deep inside draft-worker load (draft-model daemon TBD). + if self.weight_cache_mode != "off" and self.speculative_algorithm is not None: + raise ValueError( + "--weight-cache-mode is not supported together with speculative " + "decoding (--speculative-algorithm): the weight cache daemon does " + "not export the draft model's weights. Disable one of them " + "(--weight-cache-mode off) for this configuration." + ) + def _is_mistral_native_format(self) -> bool: """True iff the checkpoint requires load_format=mistral. diff --git a/python/sglang/srt/weight_cache/__init__.py b/python/sglang/srt/weight_cache/__init__.py new file mode 100644 index 000000000..4fbddca97 --- /dev/null +++ b/python/sglang/srt/weight_cache/__init__.py @@ -0,0 +1,11 @@ +# SPDX-License-Identifier: Apache-2.0 +# +# Intentionally empty: importing any weight_cache submodule (e.g. +# ``sglang.srt.weight_cache.protocol``) executes this package __init__ first. +# Eagerly re-exporting daemon/ipc_loader here would pull in torch and the model +# loader on that cheap protocol import, re-introducing the circular-import and +# startup-cost problems the local-import refactor removed. Import the concrete +# symbols from their submodules instead, e.g. +# from sglang.srt.weight_cache.protocol import CacheConfig +# from sglang.srt.weight_cache.daemon import launch_weight_cache_daemons +# from sglang.srt.weight_cache.ipc_loader import IpcModelLoader diff --git a/python/sglang/srt/weight_cache/daemon.py b/python/sglang/srt/weight_cache/daemon.py new file mode 100644 index 000000000..a47cc2643 --- /dev/null +++ b/python/sglang/srt/weight_cache/daemon.py @@ -0,0 +1,944 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Weight Cache Daemon — a persistent process that holds post-quantized, +TP-sharded model weights in GPU memory and serves them via CUDA IPC handles. + +Each GPU runs one daemon process for its TP rank. The daemon: +1. Loads model weights from disk (full pipeline: disk → TP shard → quantize) +2. Exports every parameter/buffer as a CUDA IPC handle +3. Serves handles over a Unix socket to requesting engine processes +4. Validates CacheConfig compatibility before serving + +Usage: + # Single-node: launch all TP rank daemons with a single command: + python -m sglang.srt.weight_cache.daemon \ + --model-path /path/to/model --tp-size 4 \ + --load-format auto --dtype auto --quantization fp8 + + # Multi-node: run on each node with --nnodes and --node-rank: + # Node 0: + python -m sglang.srt.weight_cache.daemon \ + --model-path /path/to/model --tp-size 16 \ + --nnodes 2 --node-rank 0 \ + --dist-init-method tcp://node0-ip:29500 + + # Node 1: + python -m sglang.srt.weight_cache.daemon \ + --model-path /path/to/model --tp-size 16 \ + --nnodes 2 --node-rank 1 \ + --dist-init-method tcp://node0-ip:29500 + + # Or launch a single daemon for a specific rank: + python -m sglang.srt.weight_cache.daemon \ + --model-path /path/to/model \ + --gpu-id 0 --tp-size 4 --tp-rank 0 \ + --dist-init-method tcp://127.0.0.1:29500 +""" + +import logging +import os +import signal +import socket +import time +from typing import Any, Dict, Optional + +import torch +import torch.distributed as dist + +from sglang.srt.configs.load_config import LoadConfig +from sglang.srt.platforms import current_platform +from sglang.srt.utils import MultiprocessingSerializer + +from .protocol import ( + CacheConfig, + check_ipc_quant_support, + cleanup_stale_daemon_files, + compute_env_stamp, + compute_global_rank, + compute_local_gpu_id, + get_quant_method_name, + get_ready_path, + get_socket_path, + hash_quant_config, + recv_msg, + send_msg, +) + +logger = logging.getLogger(__name__) + +# Per-connection timeout for the serial serve loop. A client exchange is tiny +# (a config dict + IPC handle metadata), so this generous bound never trips a +# healthy client, yet guarantees one hung/dead peer can't stall the other +# engine ranks indefinitely. +CLIENT_CONNECTION_TIMEOUT = 30.0 + + +class WeightCacheDaemon: + """Persistent GPU weight cache for a single TP rank. + + Holds the complete post-quantization state_dict in GPU memory and + serves CUDA IPC handles to engine processes via Unix socket. + """ + + def __init__( + self, + model_path: str, + gpu_id: int, + tp_size: int = 1, + tp_rank: int = 0, + pp_size: int = 1, + pp_rank: int = 0, + dp_size: int = 1, + ep_size: int = 1, + load_format: str = "auto", + dtype: str = "auto", + quantization: Optional[str] = None, + model_loader_extra_config: str = "{}", + trust_remote_code: bool = False, + revision: Optional[str] = None, + dist_init_method: Optional[str] = None, + ): + self.model_path = model_path + self.gpu_id = gpu_id + self.tp_size = tp_size + self.tp_rank = tp_rank + self.pp_size = pp_size + self.pp_rank = pp_rank + self.dp_size = dp_size + self.ep_size = ep_size + self.load_format = load_format + self.dtype = dtype + self.quantization = quantization + self.model_loader_extra_config = model_loader_extra_config + self.trust_remote_code = trust_remote_code + self.revision = revision + self.dist_init_method = dist_init_method + + self.socket_path = get_socket_path( + compute_global_rank(tp_size, pp_rank, tp_rank) + ) + self.ready_path = get_ready_path(compute_global_rank(tp_size, pp_rank, tp_rank)) + + self.model = None + self.config: Optional[CacheConfig] = None + # name -> {"handle": base64_str, "shape": list, "dtype": str, "is_param": bool} + self.state_entries: Dict[str, Dict[str, Any]] = {} + + def _init_distributed(self, server_args, model_config): + """Initialize the distributed backend required for model loading. + + Uses the same world_size/rank formula as the engine: + world_size = tp_size * pp_size + rank = tp_size * pp_rank + tp_rank + """ + from sglang.srt.distributed.parallel_state import ( + init_distributed_environment, + initialize_model_parallel, + model_parallel_is_initialized, + ) + + if model_parallel_is_initialized(): + logger.info( + f"[WeightCacheDaemon gpu={self.gpu_id}] " + f"Distributed already initialized, skipping" + ) + return + + # Initialize distributed environment + import torch.distributed as dist + + if not dist.is_initialized(): + if self.dist_init_method is None: + # Fallback: auto-assign a port. This only works for single-process. + import socket as sock_mod + + with sock_mod.socket(sock_mod.AF_INET, sock_mod.SOCK_STREAM) as s: + s.bind(("127.0.0.1", 0)) + free_port = s.getsockname()[1] + self.dist_init_method = f"tcp://127.0.0.1:{free_port}" + + init_distributed_environment( + world_size=self.tp_size * self.pp_size, + rank=compute_global_rank(self.tp_size, self.pp_rank, self.tp_rank), + distributed_init_method=self.dist_init_method, + local_rank=self.gpu_id, + backend=current_platform.get_torch_distributed_backend_str(), + ) + + initialize_model_parallel( + tensor_model_parallel_size=self.tp_size, + pipeline_model_parallel_size=self.pp_size, + expert_model_parallel_size=self.ep_size, + ) + + # Initialize DP attention state (required by some models like Qwen3 MoE) + from sglang.srt.layers.dp_attention import initialize_dp_attention + + initialize_dp_attention(server_args, model_config) + + logger.info( + f"[WeightCacheDaemon gpu={self.gpu_id} tp_rank={self.tp_rank}] " + f"Distributed backend initialized (tp_size={self.tp_size}, " + f"pp_size={self.pp_size}, " + f"world_size={self.tp_size * self.pp_size})" + ) + + def load(self): + """Full loading pipeline: disk → TP shard → quantize → export IPC handles.""" + # CUDA IPC weight sharing relies on torch's _share_cuda_ handle export, + # which only exists on CUDA-alike platforms (CUDA / ROCm). Fail loud here + # instead of dying deep inside the export with an opaque error. + if not current_platform.is_cuda_alike(): + raise RuntimeError( + f"[WeightCacheDaemon] the weight cache daemon requires a CUDA-alike " + f"platform (CUDA or ROCm) for CUDA IPC weight sharing, but the " + f"active platform device type is {current_platform.device_type!r}. " + f"Disable the weight cache (--weight-cache-mode off)." + ) + # expandable_segments makes torch's caching allocator hand out memory + # that cannot be exported via _share_cuda_, so the IPC export below would + # die mid-way with an opaque CUDA error. Fail fast with an actionable + # message before touching the device. + self._assert_ipc_compatible_allocator() + current_platform.set_device(current_platform.get_device(self.gpu_id)) + + # Reduce thread contention during multi-process loading + torch.set_num_threads(1) + + # Lazy imports to avoid circular dependencies and speed up startup + from sglang.srt.configs.device_config import DeviceConfig + from sglang.srt.configs.model_config import ModelConfig + from sglang.srt.model_loader.loader import get_model_loader + from sglang.srt.runtime_context import get_context + from sglang.srt.server_args import ServerArgs + + server_args = ServerArgs( + model_path=self.model_path, + dtype=self.dtype, + quantization=self.quantization, + trust_remote_code=self.trust_remote_code, + tp_size=self.tp_size, + pp_size=self.pp_size, + dp_size=self.dp_size, + ep_size=self.ep_size, + load_format=self.load_format, + model_loader_extra_config=self.model_loader_extra_config, + ) + get_context().set_server_args(server_args) + + # Initialize distributed backend for model loading + # (must be done after server_args and model_config are available) + # Build model config first, then init distributed + model_config = ModelConfig( + model_path=self.model_path, + trust_remote_code=self.trust_remote_code, + revision=self.revision, + dtype=self.dtype, + quantization=self.quantization, + ) + + # Build cache config fingerprint BEFORE loading the model. + # Loading may mutate hf_config.quantization_config (e.g. via + # process_weights_after_loading), which would produce a different + # hash than what the engine computes from the original config. + # ModelConfig always exposes hf_config/quantization directly; + # quantization_config is the only genuinely-optional attribute. + quant_config = getattr(model_config.hf_config, "quantization_config", None) + quant_method = get_quant_method_name( + self.quantization or model_config.quantization + ) + if not quant_method and quant_config is not None: + quant_method = get_quant_method_name(quant_config) + + self.config = CacheConfig( + model_path=self.model_path, + model_arch=( + model_config.hf_config.architectures[0] + if model_config.hf_config.architectures + else "" + ), + tp_size=self.tp_size, + tp_rank=self.tp_rank, + pp_size=self.pp_size, + pp_rank=self.pp_rank, + dp_size=self.dp_size, + ep_size=self.ep_size, + quant_method=quant_method, + quant_config_hash=hash_quant_config(quant_config), + dtype=str(model_config.dtype), + revision=self.revision or "", + **compute_env_stamp(), + ) + + # Refuse to serve quant methods not verified to round-trip through pure + # IPC tensor export. Checked before loading so an unsupported model + # fails fast instead of after minutes of disk I/O. + check_ipc_quant_support(quant_method, quant_config, where="daemon") + + # Initialize distributed backend (requires server_args + model_config) + self._init_distributed(server_args, model_config) + + # Build load config + load_config = LoadConfig( + load_format=self.load_format, + model_loader_extra_config=self.model_loader_extra_config, + tp_rank=self.tp_rank, + ) + + logger.info( + f"[WeightCacheDaemon gpu={self.gpu_id} tp_rank={self.tp_rank}] " + f"Loading model from disk: {self.model_path}" + ) + tic = time.perf_counter() + + # Load model using DefaultModelLoader (includes TP sharding + quant post-process) + loader = get_model_loader(load_config=load_config, model_config=model_config) + self.model = loader.load_model( + model_config=model_config, + device_config=DeviceConfig(current_platform.device_type, self.gpu_id), + ) + + elapsed = time.perf_counter() - tic + logger.info( + f"[WeightCacheDaemon gpu={self.gpu_id} tp_rank={self.tp_rank}] " + f"Model loaded from disk in {elapsed:.2f}s" + ) + + # Ensure every post-processing kernel has retired before we export the + # memory: clients map these tensors read-only via IPC and would otherwise + # risk observing half-written weights. + current_platform.synchronize() + + # Export all parameters and buffers as IPC handles + self._export_state() + + logger.info( + f"[WeightCacheDaemon gpu={self.gpu_id} tp_rank={self.tp_rank}] " + f"Exported {len(self.state_entries)} tensors as IPC handles. " + f"Ready to serve." + ) + + @staticmethod + def _assert_ipc_compatible_allocator() -> None: + """Reject allocator configs incompatible with CUDA IPC export. + + The expandable-segments allocator returns memory that cannot be shared + through torch's _share_cuda_ handle, which would make the export fail + partway with an opaque error. Detect it up front and fail loud. + """ + for var in ("PYTORCH_CUDA_ALLOC_CONF", "PYTORCH_ALLOC_CONF"): + conf = os.environ.get(var, "") + for field in conf.split(","): + key, _, value = field.partition(":") + if ( + key.strip() == "expandable_segments" + and value.strip().lower() == "true" + ): + raise RuntimeError( + f"[WeightCacheDaemon] {var} sets expandable_segments:True, " + f"which is incompatible with CUDA IPC weight sharing: the " + f"expandable-segments allocator hands out memory that cannot " + f"be exported via _share_cuda_, so the IPC handle export " + f"would fail mid-way. Unset expandable_segments for the " + f"weight cache daemon process (it can stay enabled for the " + f"engine itself)." + ) + + def _export_state(self): + """Export model parameters and buffers as CUDA IPC handles. + + This includes both persistent buffers (in state_dict) and non-persistent + buffers (e.g. rotary embedding cos_sin_cache) so the engine can fully + reconstruct the model state via zero-copy IPC. + """ + self.state_entries.clear() + + # remove_duplicate=False so tied weights are recognized as parameters + # under every name. state_dict() below emits both tied keys, and with the + # deduped set the duplicate would be mis-registered as a buffer, not a + # parameter, on the client. + param_names = set( + name for name, _ in self.model.named_parameters(remove_duplicate=False) + ) + state_dict_names = set(self.model.state_dict().keys()) + + # Export all items from state_dict (parameters + persistent buffers) + for name, tensor in self.model.state_dict().items(): + ipc_handle = MultiprocessingSerializer.serialize( + tensor.data, output_str=True + ) + self.state_entries[name] = { + "handle": ipc_handle, + "shape": list(tensor.shape), + "dtype": str(tensor.dtype).replace("torch.", ""), + "is_param": name in param_names, + } + + # Also export non-persistent buffers (not in state_dict but needed + # for inference, e.g. rotary embedding cos_sin_cache) + non_persistent_count = 0 + for name, buf in self.model.named_buffers(): + if name not in state_dict_names: + ipc_handle = MultiprocessingSerializer.serialize( + buf.data, output_str=True + ) + self.state_entries[name] = { + "handle": ipc_handle, + "shape": list(buf.shape), + "dtype": str(buf.dtype).replace("torch.", ""), + "is_param": False, + } + non_persistent_count += 1 + + # Log total size + total_bytes = sum( + entry["handle"].__len__() if hasattr(entry["handle"], "__len__") else 0 + for entry in self.state_entries.values() + ) + logger.info( + f"[WeightCacheDaemon gpu={self.gpu_id}] " + f"Exported {len(self.state_entries)} tensors " + f"({non_persistent_count} non-persistent buffers), " + f"serialized handle size ~{total_bytes / 1024 / 1024:.1f} MB" + ) + + def serve(self): + """Block and serve IPC handles over Unix socket.""" + # Do NOT unlink an existing socket here: stale-file cleanup is the launch + # path's job (cleanup_stale_daemon_files refuses to remove a socket whose + # .ready still points at a live PID). A leftover live socket makes bind() + # fail loudly instead of silently stealing another daemon's socket. + sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) + old_umask = os.umask(0o177) + try: + sock.bind(self.socket_path) + finally: + os.umask(old_umask) + sock.listen(8) + sock.settimeout(1.0) # Allow periodic shutdown check + + # Write ready file + with open(self.ready_path, "w") as f: + f.write(f"pid={os.getpid()}\n") + f.write(f"config={self.config.to_dict()}\n") + + logger.info( + f"[WeightCacheDaemon gpu={self.gpu_id}] " f"Listening on {self.socket_path}" + ) + + self._running = True + + def _signal_handler(signum, frame): + logger.info( + f"[WeightCacheDaemon gpu={self.gpu_id}] Received signal {signum}, shutting down" + ) + self._running = False + + signal.signal(signal.SIGTERM, _signal_handler) + signal.signal(signal.SIGINT, _signal_handler) + + try: + while self._running: + try: + conn, _ = sock.accept() + # The listen-socket timeout above only bounds accept(); the + # accepted connection is blocking by default. Since we serve + # connections serially, a client that connects but never sends + # (or dies mid-send) would block recv_msg forever and stall + # every other engine rank. Bound each exchange instead. + conn.settimeout(CLIENT_CONNECTION_TIMEOUT) + try: + self._handle_connection(conn) + except Exception as e: + logger.error( + f"[WeightCacheDaemon gpu={self.gpu_id}] " + f"Error handling connection: {e}", + exc_info=True, + ) + finally: + conn.close() + except socket.timeout: + continue + finally: + sock.close() + if os.path.exists(self.socket_path): + os.unlink(self.socket_path) + if os.path.exists(self.ready_path): + os.unlink(self.ready_path) + logger.info(f"[WeightCacheDaemon gpu={self.gpu_id}] Shutdown complete") + + def _handle_connection(self, conn: socket.socket): + """Handle a single client connection.""" + req = recv_msg(conn) + + if req.get("type") == "query_config": + # Client asks for config without requesting handles + send_msg(conn, {"status": "ok", "config": self.config.to_dict()}) + + elif req.get("type") == "fetch_state": + # Client requests full state with IPC handles + engine_config = CacheConfig.from_dict(req["config"]) + if not self.config.matches(engine_config): + # Log detailed mismatch info for debugging + daemon_dict = self.config.to_dict() + engine_dict = engine_config.to_dict() + mismatches = { + k: (daemon_dict.get(k), engine_dict.get(k)) + for k in daemon_dict + if daemon_dict.get(k) != engine_dict.get(k) + } + logger.warning( + f"[WeightCacheDaemon gpu={self.gpu_id}] " + f"Config mismatch: {mismatches}" + ) + send_msg( + conn, {"status": "mismatch", "daemon_config": self.config.to_dict()} + ) + return + + logger.info( + f"[WeightCacheDaemon gpu={self.gpu_id}] " + f"Serving {len(self.state_entries)} IPC handles to engine" + ) + send_msg( + conn, + { + "status": "ok", + "config": self.config.to_dict(), + "entries": self.state_entries, + # PID so the client can watch daemon liveness: if this + # process dies while clients hold IPC mappings, their + # param.data (and any CUDA-graph-captured addresses) dangle. + "pid": os.getpid(), + }, + ) + + elif req.get("type") == "ping": + send_msg(conn, {"status": "ok"}) + + else: + send_msg( + conn, + { + "status": "error", + "message": f"Unknown request type: {req.get('type')}", + }, + ) + + def shutdown(self): + """Release GPU memory and clean up.""" + if dist.is_initialized(): + dist.destroy_process_group() + if self.model is not None: + del self.model + self.model = None + self.state_entries.clear() + current_platform.empty_cache() + self._running = False + + +def run_weight_cache_daemon( + model_path: str, + gpu_id: int, + tp_size: int = 1, + tp_rank: int = 0, + pp_size: int = 1, + pp_rank: int = 0, + dp_size: int = 1, + ep_size: int = 1, + load_format: str = "auto", + dtype: str = "auto", + quantization: Optional[str] = None, + model_loader_extra_config: str = "{}", + trust_remote_code: bool = False, + revision: Optional[str] = None, + dist_init_method: Optional[str] = None, +): + """Entry point for running a weight cache daemon process.""" + logging.basicConfig( + level=logging.INFO, + format=f"%(asctime)s [Daemon gpu={gpu_id} tp_rank={tp_rank}] %(levelname)s %(message)s", + ) + + # Die if our parent (the engine or the standalone launcher that spawned us) + # dies, even on SIGKILL/OOM-kill. Without this an orphaned daemon keeps a + # full weight copy pinned in GPU memory and its live-PID .ready file blocks + # the next launch — the opposite of fast recovery. + from sglang.srt.utils import kill_itself_when_parent_died + + kill_itself_when_parent_died() + + daemon = WeightCacheDaemon( + model_path=model_path, + gpu_id=gpu_id, + tp_size=tp_size, + tp_rank=tp_rank, + pp_size=pp_size, + pp_rank=pp_rank, + dp_size=dp_size, + ep_size=ep_size, + load_format=load_format, + dtype=dtype, + quantization=quantization, + model_loader_extra_config=model_loader_extra_config, + trust_remote_code=trust_remote_code, + revision=revision, + dist_init_method=dist_init_method, + ) + + daemon.load() + daemon.serve() + + +def launch_weight_cache_daemons( + model_path: str, + tp_size: int = 1, + pp_size: int = 1, + dp_size: int = 1, + ep_size: int = 1, + nnodes: int = 1, + node_rank: int = 0, + base_gpu_id: int = 0, + gpu_id_step: int = 1, + load_format: str = "auto", + dtype: str = "auto", + quantization: Optional[str] = None, + model_loader_extra_config: str = "{}", + trust_remote_code: bool = False, + revision: Optional[str] = None, + dist_init_method: Optional[str] = None, + timeout: int = 1800, + force: bool = False, +): + """Launch weight cache daemon processes for this node's PP×TP ranks. + + For single-node (nnodes=1): spawns pp_size * tp_size daemons. + For multi-node (nnodes>1): spawns this node's share of PP×TP daemons, + mapping local gpu_id to the correct global (pp_rank, tp_rank). + + Uses subprocess.Popen instead of multiprocessing.Process to avoid + initializing CUDA in the parent process, which can degrade CUDA IPC + performance in child processes. + + Usage (single-node): + python -m sglang.srt.weight_cache.daemon \\ + --model-path /path/to/model --tp-size 4 + + Usage (multi-node, run on each node): + # Node 0: + python -m sglang.srt.weight_cache.daemon \\ + --model-path /path/to/model --tp-size 16 \\ + --nnodes 2 --node-rank 0 \\ + --dist-init-method tcp://node0-ip:29500 + + # Node 1: + python -m sglang.srt.weight_cache.daemon \\ + --model-path /path/to/model --tp-size 16 \\ + --nnodes 2 --node-rank 1 \\ + --dist-init-method tcp://node0-ip:29500 + """ + import socket as sock_mod + import subprocess + import sys + + # Replicate _calculate_rank_ranges logic from engine.py + pp_size_per_node = max(pp_size // nnodes, 1) + nnodes_per_pp_rank = max(nnodes // pp_size, 1) + pp_rank_range = range( + pp_size_per_node * (node_rank // nnodes_per_pp_rank), + pp_size_per_node * (node_rank // nnodes_per_pp_rank + 1), + ) + nnodes_per_tp_group = nnodes_per_pp_rank + tp_size_per_node = tp_size // nnodes_per_tp_group + tp_rank_range = range( + tp_size_per_node * (node_rank % nnodes_per_tp_group), + tp_size_per_node * (node_rank % nnodes_per_tp_group + 1), + ) + + if nnodes > 1 and dist_init_method is None: + raise ValueError( + "dist_init_method is required for multi-node weight cache daemons. " + "Use --dist-init-method tcp://: to specify the " + "rendezvous address accessible from all nodes." + ) + + # Auto-allocate a free port for the distributed init method + if dist_init_method is None: + with sock_mod.socket(sock_mod.AF_INET, sock_mod.SOCK_STREAM) as s: + s.bind(("127.0.0.1", 0)) + free_port = s.getsockname()[1] + dist_init_method = f"tcp://127.0.0.1:{free_port}" + + python_path = sys.executable + daemon_module = "sglang.srt.weight_cache.daemon" + + # 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(tp_size, pp_rank, tp_rank) + cleanup_stale_daemon_files(global_rank, force=force) + + procs = [] + for pp_rank in pp_rank_range: + for tp_rank in tp_rank_range: + gpu_id = compute_local_gpu_id( + pp_rank, + tp_rank, + pp_size_per_node, + tp_size_per_node, + base_gpu_id=base_gpu_id, + gpu_id_step=gpu_id_step, + ) + cmd = [ + python_path, + "-m", + daemon_module, + "--model-path", + model_path, + "--gpu-id", + str(gpu_id), + "--tp-size", + str(tp_size), + "--tp-rank", + str(tp_rank), + "--dp-size", + str(dp_size), + "--ep-size", + str(ep_size), + "--pp-size", + str(pp_size), + "--pp-rank", + str(pp_rank), + "--load-format", + load_format, + "--dtype", + dtype, + "--dist-init-method", + dist_init_method, + ] + if quantization: + cmd += ["--quantization", quantization] + if model_loader_extra_config and model_loader_extra_config != "{}": + cmd += ["--model-loader-extra-config", model_loader_extra_config] + if trust_remote_code: + cmd += ["--trust-remote-code"] + if revision: + cmd += ["--revision", revision] + + proc = subprocess.Popen(cmd) + procs.append(proc) + logger.info( + f"Launched weight cache daemon gpu={gpu_id} " + f"pp_rank={pp_rank} tp_rank={tp_rank} pid={proc.pid}" + ) + + # Wait for all daemons on this node to become ready + num_daemons = len(procs) + check_interval = 2 + start_time = time.time() + 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) + while not os.path.exists(ready_path): + time.sleep(check_interval) + if time.time() - start_time > timeout: + logger.error( + f"Weight cache daemon pp_rank={pp_rank} tp_rank={tp_rank} " + f"did not become ready within {timeout}s" + ) + for p in procs: + p.terminate() + raise TimeoutError( + f"Weight cache daemon pp_rank={pp_rank} tp_rank={tp_rank} " + f"did not become ready within {timeout}s" + ) + # Check if any daemon exited prematurely + for p in procs: + retcode = p.poll() + if retcode is not None: + logger.error( + f"Weight cache daemon exited prematurely " + f"with code {retcode}" + ) + for other in procs: + if other.poll() is None: + other.terminate() + raise RuntimeError( + f"Weight cache daemon exited prematurely " + f"with code {retcode}" + ) + logger.info( + f"Weight cache daemon pp_rank={pp_rank} tp_rank={tp_rank} is ready" + ) + + logger.info( + f"All {num_daemons} weight cache daemons on node {node_rank} are ready " + f"(pp_ranks={pp_rank_range.start}..{pp_rank_range.stop - 1}, " + f"tp_ranks={tp_rank_range.start}..{tp_rank_range.stop - 1}, " + f"dist_init_method={dist_init_method})" + ) + + # Monitor daemons — poll all of them and, the moment any one exits, + # terminate the rest and raise. A serial proc.wait() would not notice a + # mid-list death (e.g. procs[1] dying while procs[0] is still alive) until + # the earlier proc happened to exit, and it never surfaced the failure. + exited = None + try: + while exited is None: + for proc in procs: + if proc.poll() is not None: + exited = proc + break + else: + time.sleep(1) + continue + except KeyboardInterrupt: + logger.info("Received KeyboardInterrupt, shutting down daemons") + finally: + for proc in procs: + if proc.poll() is None: + proc.terminate() + for proc in procs: + try: + proc.wait(timeout=5) + except subprocess.TimeoutExpired: + proc.kill() + logger.info("All weight cache daemons have been terminated") + + if exited is not None: + raise RuntimeError( + f"Weight cache daemon (pid={exited.pid}) exited with code " + f"{exited.returncode}; terminated the remaining daemons." + ) + + +if __name__ == "__main__": + import argparse + + parser = argparse.ArgumentParser(description="SGLang Weight Cache Daemon") + parser.add_argument("--model-path", required=True, help="Path to model weights") + parser.add_argument("--tp-size", type=int, default=1, help="Tensor parallel size") + parser.add_argument( + "--gpu-id", + type=int, + default=None, + help="GPU device ID for a single daemon. " + "If omitted, launches daemons for all TP ranks (0..tp_size-1).", + ) + parser.add_argument( + "--tp-rank", + type=int, + default=None, + help="TP rank for a single daemon. " + "If omitted, launches daemons for all TP ranks.", + ) + parser.add_argument("--dp-size", type=int, default=1, help="Data parallel size") + parser.add_argument("--ep-size", type=int, default=1, help="Expert parallel size") + parser.add_argument("--pp-size", type=int, default=1, help="Pipeline parallel size") + parser.add_argument("--pp-rank", type=int, default=0, help="Pipeline parallel rank") + parser.add_argument("--nnodes", type=int, default=1, help="Total number of nodes") + parser.add_argument( + "--base-gpu-id", + type=int, + default=0, + help="GPU id of this node's first rank (mirrors the engine's " + "--base-gpu-id). Used to place daemons on the same GPUs the engine " + "ranks will use.", + ) + parser.add_argument( + "--gpu-id-step", + type=int, + default=1, + help="Stride between consecutive ranks' GPU ids (mirrors the engine's " + "--gpu-id-step).", + ) + parser.add_argument( + "--node-rank", + type=int, + default=0, + help="Rank of this node (0-indexed). Required for multi-node.", + ) + parser.add_argument("--load-format", default="auto", help="Weight load format") + parser.add_argument("--dtype", default="auto", help="Model dtype") + parser.add_argument("--quantization", default=None, help="Quantization method") + parser.add_argument( + "--model-loader-extra-config", + default="{}", + help="Extra config for model loader (JSON string)", + ) + parser.add_argument("--trust-remote-code", action="store_true") + parser.add_argument("--revision", default=None, help="Model revision") + parser.add_argument( + "--dist-init-method", + default=None, + help="Distributed init method (e.g. tcp://node0-ip:29500). " + "Auto-assigned for single-node when launching all ranks. " + "Required for multi-node (nnodes > 1) and must be accessible " + "from all nodes. Also required for tp_size > 1 when launching " + "a single rank.", + ) + parser.add_argument( + "--timeout", + type=int, + default=1800, + help="Timeout in seconds to wait for all daemons to become ready (default: 1800)", + ) + parser.add_argument( + "--force", + action="store_true", + help="Take over a rank whose .ready file still points at a live PID by " + "killing that daemon (use to reclaim a wedged/orphaned daemon).", + ) + + args = parser.parse_args() + + if args.gpu_id is not None or args.tp_rank is not None: + # Single-rank mode: launch one daemon for the specified rank + gpu_id = args.gpu_id if args.gpu_id is not None else args.tp_rank + tp_rank = args.tp_rank if args.tp_rank is not None else args.gpu_id + # Refuse to clobber a live daemon already holding this rank (mirrors the + # multi-rank launcher and the engine path, which the multi-rank spawns + # of this same entrypoint rely on). --force kills and takes over. + cleanup_stale_daemon_files( + compute_global_rank(args.tp_size, args.pp_rank, tp_rank), + force=args.force, + ) + run_weight_cache_daemon( + model_path=args.model_path, + gpu_id=gpu_id, + tp_size=args.tp_size, + tp_rank=tp_rank, + pp_size=args.pp_size, + pp_rank=args.pp_rank, + dp_size=args.dp_size, + ep_size=args.ep_size, + load_format=args.load_format, + dtype=args.dtype, + quantization=args.quantization, + model_loader_extra_config=args.model_loader_extra_config, + trust_remote_code=args.trust_remote_code, + revision=args.revision, + dist_init_method=args.dist_init_method, + ) + else: + # Multi-rank mode: launch daemons for this node's TP ranks + launch_weight_cache_daemons( + model_path=args.model_path, + tp_size=args.tp_size, + pp_size=args.pp_size, + dp_size=args.dp_size, + ep_size=args.ep_size, + nnodes=args.nnodes, + node_rank=args.node_rank, + base_gpu_id=args.base_gpu_id, + gpu_id_step=args.gpu_id_step, + load_format=args.load_format, + dtype=args.dtype, + quantization=args.quantization, + model_loader_extra_config=args.model_loader_extra_config, + trust_remote_code=args.trust_remote_code, + revision=args.revision, + dist_init_method=args.dist_init_method, + timeout=args.timeout, + force=args.force, + ) diff --git a/python/sglang/srt/weight_cache/ipc_loader.py b/python/sglang/srt/weight_cache/ipc_loader.py new file mode 100644 index 000000000..02e7ae772 --- /dev/null +++ b/python/sglang/srt/weight_cache/ipc_loader.py @@ -0,0 +1,565 @@ +# SPDX-License-Identifier: Apache-2.0 +"""IPC Model Loader — loads model weights from a Weight Cache Daemon via CUDA IPC. + +Zero-copy mode: param.data points directly to IPC-mapped GPU memory. Only 1x GPU +memory needed — engine and daemon share the same physical GPU memory via CUDA IPC. +Engine depends on daemon staying alive. +""" + +import logging +import os +import signal +import stat +import threading +import time +from typing import Optional + +import torch +import torch.nn as nn + +from sglang.srt.configs.load_config import LoadConfig +from sglang.srt.model_loader.loader import ( + BaseModelLoader, + _initialize_model, +) +from sglang.srt.utils import MultiprocessingSerializer + +from .protocol import ( + CacheConfig, + check_ipc_quant_support, + compute_env_stamp, + get_quant_method_name, + hash_quant_config, + recv_msg, + send_msg, +) + +logger = logging.getLogger(__name__) + +# How often the client polls the serving daemon's PID for liveness. +_DAEMON_LIVENESS_POLL_INTERVAL = 5.0 + + +class IpcModelLoader(BaseModelLoader): + """Load model weights from a Weight Cache Daemon via CUDA IPC handles. + + In daemon mode (weight_cache_mode="daemon"), the engine and daemon share + the same GPU. Falling back to disk loading would cause OOM because both + processes would hold weights on the same GPU. Therefore, daemon mode + raises an error if the daemon is unavailable instead of falling back. + + In client mode, disk fallback is allowed ONLY when the daemon is genuinely + absent (its Unix socket file does not exist). Every other failure is a hard + error rather than a silent fallback, so a broken IPC path never masquerades + as a healthy (but slow, disk-loaded) server: + + - socket file missing -> fall back to disk load + - connection refused -> raise (daemon crashed after binding) + - CacheConfig mismatch -> raise (do NOT disk-load on a shared GPU + holding a different config's weights; + also surfaces fingerprint drift bugs) + - any protocol / transfer error -> raise + + See _fetch_from_cache for the authoritative fallback-vs-raise contract. + """ + + def __init__( + self, + load_config: LoadConfig, + socket_path: str, + fallback_loader_cls=None, + weight_cache_mode: str = "client", + fallback_load_format: str = "auto", + ): + super().__init__(load_config) + self.socket_path = socket_path + self.weight_cache_mode = weight_cache_mode + self._fallback_loader_cls = fallback_loader_cls + self._fallback_load_format = fallback_load_format + + def load_model( + self, + *, + model_config, + device_config, + ) -> nn.Module: + """Load model weights from the weight cache daemon. + + In daemon mode, raises RuntimeError if the daemon is unavailable + (fallback to disk loading would cause OOM on shared GPUs). + In client mode, falls back to DefaultModelLoader. + """ + tic = time.perf_counter() + + # Hard-gate unsupported quant methods before touching the daemon, so an + # unsupported model fails explicitly instead of silently disk-loading + # (client mode) or serving wrong-numerics IPC weights. Checked here so + # it applies regardless of whether the daemon is reachable. + quant_method, engine_quant_config = self._resolve_engine_quant(model_config) + 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) + + if cache_data is None: + if self.weight_cache_mode == "daemon": + raise RuntimeError( + f"[IpcModelLoader] Weight cache daemon not available at " + f"{self.socket_path}. In daemon mode, fallback to disk " + f"loading is disabled because the daemon process already " + f"holds weights on the same GPU — loading from disk would " + f"cause OOM. Please ensure the weight cache daemon is " + f"running and the config matches." + ) + logger.warning( + "[IpcModelLoader] Weight cache not available or config mismatch, " + "falling back to disk load" + ) + return self._fallback_load(model_config, device_config) + + entries = cache_data["entries"] + logger.info( + f"[IpcModelLoader] Fetched {len(entries)} IPC handles from daemon " + f"in {time.perf_counter() - tic:.2f}s" + ) + + from sglang.srt.model_loader.loader import ( + _get_quantization_config, + ) + + quant_config = _get_quantization_config(model_config, self.load_config) + + model = self._load_zero_copy_mode( + model_config, + device_config, + entries, + quant_config, + ) + + # Skip _post_load_weights: the daemon already ran + # process_weights_after_loading on the weights before exporting + # IPC handles. Running it again would double-process (e.g., + # re-quantize already-quantized weights), corrupting tensor data. + + # Rebuild stale tensor views. Some modules store tensor views as + # plain attributes (not parameters/buffers) during __init__. When + # the model is initialized on meta device and then weights are + # replaced via IPC mapping, these views still point to the old + # meta storage. We must recreate them from the now-valid tensors. + self._rebuild_stale_views(model) + + # The model now points into the daemon's GPU memory via CUDA IPC. If the + # daemon dies, those pointers dangle, so watch it and fail loud. + self._start_daemon_liveness_watchdog(cache_data.get("pid")) + + logger.info( + f"[IpcModelLoader] Loaded model via IPC (mode={self.weight_cache_mode}), " + f"total={time.perf_counter() - tic:.2f}s" + ) + + return model.eval() + + def _start_daemon_liveness_watchdog(self, daemon_pid: Optional[int]) -> None: + """Fail loud if the serving daemon dies while we hold its weights. + + In both client and (engine-spawned) daemon mode, the model's param.data + points into the daemon's GPU memory via CUDA IPC, and CUDA graphs may + capture those addresses. If the daemon exits, the pointers dangle: + forward passes would read freed GPU memory -> illegal-address crashes or + silent garbage. There is no safe in-place recovery, so a background + thread polls the daemon PID and, on death, SIGKILLs this process with a + clear message instead of letting it serve corrupt results. + """ + if not daemon_pid or daemon_pid <= 0: + logger.warning( + "[IpcModelLoader] Daemon did not report a PID; skipping the " + "daemon-liveness watchdog. A daemon crash will not be detected." + ) + return + + def _daemon_alive(pid: int) -> bool: + try: + os.kill(pid, 0) + except ProcessLookupError: + return False + except PermissionError: + return True # exists but owned by another user + return True + + def _watch() -> None: + while True: + time.sleep(_DAEMON_LIVENESS_POLL_INTERVAL) + if not _daemon_alive(daemon_pid): + logger.critical( + f"[IpcModelLoader] Weight cache daemon (pid={daemon_pid}) " + f"died while this engine holds its weights via CUDA IPC. " + f"The mapped weight pointers are now dangling; continuing " + f"would read freed GPU memory. Terminating this process." + ) + os.kill(os.getpid(), signal.SIGKILL) + return + + threading.Thread( + target=_watch, name="weight-cache-daemon-watchdog", daemon=True + ).start() + logger.info( + f"[IpcModelLoader] Started daemon-liveness watchdog for pid={daemon_pid}" + ) + + def _resolve_engine_quant(self, model_config): + """Return (quant_method, quant_config) matching the daemon's fingerprint. + + Shared by the IPC allowlist gate and the CacheConfig fingerprint so the + two can never drift apart. ModelConfig always exposes + hf_config/quantization directly; quantization_config is the only + genuinely-optional attribute. + """ + quant_config = getattr(model_config.hf_config, "quantization_config", None) + quant_method = get_quant_method_name(model_config.quantization) + if not quant_method and quant_config is not None: + quant_method = get_quant_method_name(quant_config) + return quant_method, quant_config + + @staticmethod + def _rebuild_stale_views(model): + """Rebuild tensor views that went stale after IPC weight replacement. + + RadixLinearAttention.conv_weights is a view of conv1d.weight created + during __init__. After IPC mapping replaces conv1d.weight with a new + tensor, the old view still points to meta-device storage. Recreate + it from the now-valid parameter. + """ + try: + from sglang.srt.layers.radix_linear_attention import RadixLinearAttention + except ImportError: + return + + count = 0 + for _, module in model.named_modules(): + conv1d = getattr(module, "conv1d", None) + attn = getattr(module, "attn", None) + if conv1d is not None and isinstance(attn, RadixLinearAttention): + if hasattr(conv1d, "weight") and conv1d.weight is not None: + attn.conv_weights = conv1d.weight.view( + conv1d.weight.size(0), conv1d.weight.size(2) + ) + if hasattr(conv1d, "bias") and conv1d.bias is not None: + attn.bias = conv1d.bias + count += 1 + + if count > 0: + logger.info(f"[IpcModelLoader] Rebuilt {count} stale conv_weights views") + + @staticmethod + def _set_module_tensor(model, name, tensor, is_param=True): + """Replace or register a parameter/buffer in the model by its full dotted name. + + This is necessary because setting param.data on a meta-device tensor + raises a type mismatch error (meta and CUDA tensors have incompatible + dispatch keys). Instead, we walk the module tree and use setattr to + replace the entire parameter/buffer object. + + If the attribute already exists as a parameter/buffer, it is replaced. + If it doesn't exist (e.g. post-quantization params like weight_scale), + it is registered as a new parameter or buffer. + """ + parts = name.split(".") + obj = model + for part in parts[:-1]: + obj = getattr(obj, part) + leaf_name = parts[-1] + if is_param: + # requires_grad=False: the IPC memory is shared/read-only and SGLang + # is inference-only, so autograd must never write into it. + new_param = nn.Parameter(tensor, requires_grad=False) + setattr(obj, leaf_name, new_param) + else: + # register_buffer raises KeyError if the name already exists as a + # parameter or plain attribute (not a buffer). This happens when + # process_weights_after_loading converts a parameter to a buffer + # (e.g. Mamba's A_log). Remove the old attribute first. + if leaf_name in obj._parameters: + del obj._parameters[leaf_name] + elif hasattr(obj, leaf_name) and leaf_name not in obj._buffers: + delattr(obj, leaf_name) + obj.register_buffer(leaf_name, tensor) + + def _load_zero_copy_mode( + self, + model_config, + device_config, + entries, + quant_config, + ) -> nn.Module: + """Zero-copy load: map IPC tensors directly as param.data. + + The model is initialized on the meta device (no memory allocation), + then each parameter's data is replaced with the IPC-mapped GPU tensor. + The engine and daemon share the same physical GPU memory via CUDA IPC. + """ + from sglang.srt.model_loader.utils import set_default_torch_dtype + + # Initialize model on meta device to avoid any GPU/CPU memory allocation. + # This creates the model structure with the correct parameter shapes/dtypes + # but without allocating actual storage. + with set_default_torch_dtype(model_config.dtype): + with torch.device("meta"): + model = _initialize_model( + model_config, + self.load_config, + quant_config, + ) + + # Build lookup dicts of existing parameter/buffer names in the + # meta-device model. Post-quantization parameters (e.g. weight_scale + # from FP8) are created by process_weights_after_loading, which the + # daemon already ran. These params exist in the daemon's entries but + # NOT in the meta-device model — we must register them as new attrs. + # Use dicts (not sets) so we can do O(1) shape/dtype validation + # without re-traversing the model tree on every lookup. + # remove_duplicate=False mirrors the daemon's export (which keys tied + # weights under every name) so a tied parameter is recognized under all + # of its names here too. + existing_params = { + name: param + for name, param in model.named_parameters(remove_duplicate=False) + } + existing_buffers = {name: buf for name, buf in model.named_buffers()} + existing_names = set(existing_params) | set(existing_buffers) + + imported_refs = [] + imported_count = 0 + mismatched = [] + new_params_count = 0 + map_tic = time.perf_counter() + + # Iterate over ALL daemon entries (not just model params/buffers). + # This ensures post-quantization parameters (weight_scale, etc.) + # that were created by process_weights_after_loading are also mapped. + for name, entry in entries.items(): + imported_tensor = MultiprocessingSerializer.deserialize(entry["handle"]) + is_param = entry.get("is_param", True) + + if name in existing_names: + # Existing parameter/buffer — validate shape/dtype + if name in existing_params: + ref_param = existing_params[name] + else: + ref_param = existing_buffers[name] + if ( + imported_tensor.shape != ref_param.shape + or imported_tensor.dtype != ref_param.dtype + ): + mismatched.append( + f" {name}: IPC={imported_tensor.shape}/{imported_tensor.dtype} " + f"vs model={ref_param.shape}/{ref_param.dtype}" + ) + del imported_tensor + continue + + # Replace or register the tensor in the model + self._set_module_tensor(model, name, imported_tensor, is_param=is_param) + imported_refs.append(imported_tensor) + imported_count += 1 + + if name not in existing_names: + new_params_count += 1 + + if mismatched: + raise RuntimeError( + f"[IpcModelLoader] {len(mismatched)} tensor(s) have shape/dtype " + f"mismatch between the IPC daemon and the meta-initialized model. " + f"The quantization method passed the IPC allowlist gate " + f"(check_ipc_quant_support), so this is NOT an unsupported-quant " + f"case — it indicates the daemon's weight fingerprint is " + f"incomplete or the daemon/client configs drifted (a bug to fix), " + f"not merely uninitialized weights:\n" + "\n".join(mismatched) + ) + + # After mapping every daemon entry, any tensor still on the meta device + # is one the daemon did NOT provide. Filling it with torch.empty() would + # hand the model uninitialized GPU memory — silently producing wrong + # output, the worst failure mode for a load path. Hard-error and list the + # offenders instead. + # + # The daemon exports the full state_dict AND non-persistent buffers + # (e.g. rotary embedding cos_sin_cache), so a correct setup leaves nothing + # on meta here. A non-empty list means the daemon's export is incomplete, + # or the model has a genuinely-recomputable buffer that must be recomputed + # explicitly (not filled with garbage) — add that handling here if needed. + still_on_meta_params = [ + name + for name, param in model.named_parameters() + if param.device.type == "meta" + ] + still_on_meta_buffers = [ + name for name, buf in model.named_buffers() if buf.device.type == "meta" + ] + + if still_on_meta_params or still_on_meta_buffers: + raise RuntimeError( + f"[IpcModelLoader] After IPC mapping, " + f"{len(still_on_meta_params)} parameter(s) and " + f"{len(still_on_meta_buffers)} buffer(s) remain on the meta device " + f"— the daemon did not export them. Refusing to fill them with " + f"uninitialized memory, which would silently produce wrong output. " + f"This means the daemon's export is incomplete, or a recomputable " + f"buffer needs explicit recompute logic here.\n" + f" params: {still_on_meta_params[:10]}" + f"{'...' if len(still_on_meta_params) > 10 else ''}\n" + f" buffers: {still_on_meta_buffers[:10]}" + f"{'...' if len(still_on_meta_buffers) > 10 else ''}" + ) + + map_elapsed = time.perf_counter() - map_tic + + # Stash IPC refs on the model to prevent GC (which would unmap the memory) + if imported_refs: + model._ipc_imported_tensors = imported_refs + + logger.info( + f"[IpcModelLoader] Zero-copy: mapped {imported_count} tensors " + f"({new_params_count} new post-quant), time={map_elapsed:.3f}s" + ) + + return model + + def _fetch_from_cache(self, model_config) -> Optional[dict]: + """Connect to daemon, validate config, fetch IPC handles. + + Returns the daemon response dict on success, None if the daemon is + genuinely absent (socket file doesn't exist). Raises on all other + failures so they are never silently swallowed as a disk-load fallback. + """ + import socket as socket_mod + + # 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). + try: + st = os.lstat(self.socket_path) + except FileNotFoundError: + logger.info( + f"[IpcModelLoader] Daemon socket not found at {self.socket_path}." + ) + return None + if not stat.S_ISSOCK(st.st_mode) or st.st_uid != os.getuid(): + raise RuntimeError( + f"[IpcModelLoader] Refusing to connect: {self.socket_path} is not " + f"a socket owned by this user." + ) + + sock = socket_mod.socket(socket_mod.AF_UNIX, socket_mod.SOCK_STREAM) + try: + sock.settimeout(30) + sock.connect(self.socket_path) + except FileNotFoundError: + # Raced: socket removed between lstat and connect -> treat as absent. + sock.close() + return None + except ConnectionRefusedError: + sock.close() + raise RuntimeError( + f"[IpcModelLoader] Daemon socket exists at {self.socket_path} but " + f"refused the connection. The daemon may have crashed after " + f"creating the socket. Check daemon logs." + ) + except Exception as e: + sock.close() + raise RuntimeError( + f"[IpcModelLoader] Failed to connect to daemon at " + f"{self.socket_path}: {e}" + ) from e + + try: + # Build engine's config fingerprint + from sglang.srt.runtime_context import get_parallel + + ps = get_parallel() + tp_size = ps.tp_size + tp_rank = ps.tp_rank + + pp_size = ps.pp_size + pp_rank = ps.pp_rank + + ep_size = ps.moe_ep_size + + from sglang.srt.runtime_context import get_server_args + + dp_size = get_server_args().dp_size + + quant_method, quant_config = self._resolve_engine_quant(model_config) + + engine_config = CacheConfig( + model_path=model_config.model_path, + model_arch=( + model_config.hf_config.architectures[0] + if model_config.hf_config.architectures + else "" + ), + tp_size=tp_size, + tp_rank=tp_rank, + pp_size=pp_size, + pp_rank=pp_rank, + dp_size=dp_size, + ep_size=ep_size, + quant_method=quant_method, + quant_config_hash=hash_quant_config(quant_config), + dtype=str(model_config.dtype), + revision=model_config.revision or "", + **compute_env_stamp(), + ) + + logger.info( + f"[IpcModelLoader] Requesting weights from daemon at " + f"{self.socket_path} with config: " + f"model={engine_config.model_path}, " + f"arch={engine_config.model_arch}, " + f"tp={engine_config.tp_size}/{engine_config.tp_rank}, " + f"quant={engine_config.quant_method}, " + f"dtype={engine_config.dtype}" + ) + + send_msg(sock, {"type": "fetch_state", "config": engine_config.to_dict()}) + result = recv_msg(sock) + + if result.get("status") != "ok": + daemon_config = result.get("daemon_config", {}) + raise RuntimeError( + f"[IpcModelLoader] Daemon config mismatch!\n" + f" Engine config: {engine_config.to_dict()}\n" + f" Daemon config: {daemon_config}" + ) + + return result + + except RuntimeError: + raise + except Exception as e: + raise RuntimeError( + f"[IpcModelLoader] Error communicating with daemon at " + f"{self.socket_path}: {e}" + ) from e + finally: + sock.close() + + def _fallback_load(self, model_config, device_config) -> nn.Module: + """Fall back to DefaultModelLoader for disk-based loading.""" + from sglang.srt.configs.load_config import LoadConfig + from sglang.srt.model_loader.loader import DefaultModelLoader + + fallback_config = LoadConfig( + load_format=self._fallback_load_format, + download_dir=self.load_config.download_dir, + model_loader_extra_config=self.load_config.model_loader_extra_config, + tp_rank=self.load_config.tp_rank, + ) + loader_cls = self._fallback_loader_cls or DefaultModelLoader + fallback = loader_cls(fallback_config) + return fallback.load_model( + model_config=model_config, device_config=device_config + ) + + def download_model(self, model_config) -> None: + """No-op: daemon handles its own model downloading.""" + pass diff --git a/python/sglang/srt/weight_cache/protocol.py b/python/sglang/srt/weight_cache/protocol.py new file mode 100644 index 000000000..8b3b4b341 --- /dev/null +++ b/python/sglang/srt/weight_cache/protocol.py @@ -0,0 +1,380 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Protocol definitions for the weight cache daemon. + +Defines CacheConfig for validation and socket message protocol helpers. +""" + +import hashlib +import json +import logging +import os +import pickle +import signal +import struct +from typing import Any, Dict, Optional + +import msgspec + +from sglang.srt.utils.common import safe_pickle_loads + +logger = logging.getLogger(__name__) + +# Socket path template for weight cache daemons (keyed by global rank +# = tp_size * pp_rank + tp_rank, so multi-node / multi-PP don't collide) +WEIGHT_CACHE_SOCKET_TEMPLATE = "/tmp/sglang_weight_cache_rank{global_rank}.sock" + +# Ready file template — daemon writes this after loading completes +WEIGHT_CACHE_READY_TEMPLATE = "/tmp/sglang_weight_cache_rank{global_rank}.ready" + + +class CacheConfig(msgspec.Struct): + """Fingerprint of the cached weights. Used to validate compatibility + between a daemon's cached state and a requesting engine process. + + Any mismatch triggers a fallback to disk loading. + """ + + model_path: str + model_arch: str + tp_size: int + tp_rank: int + pp_size: int + pp_rank: int + dp_size: int + ep_size: int + quant_method: str # e.g. "fp8", "gptq_marlin", "" for unquantized + quant_config_hash: str # SHA-256 hash of quantization config + dtype: str # e.g. "torch.float16" + revision: str # model revision the weights were loaded from ("" if unset) + # Environment stamp: a daemon and a client that ran different post-processing + # branches (different GPU compute capability or torch/kernel version) can + # produce incompatible weights that would map cleanly yet serve garbage. + # Comparing these turns that into a clean mismatch. See compute_env_stamp(). + device_capability: str # local compute capability, e.g. "8.0" ("" if N/A) + torch_version: str # torch.__version__ of the process that built the weights + + def matches(self, other: "CacheConfig") -> bool: + """Check if two configs are compatible for weight sharing.""" + return self == other + + def to_dict(self) -> Dict[str, Any]: + return {f: getattr(self, f) for f in self.__struct_fields__} + + @classmethod + def from_dict(cls, d: Dict[str, Any]) -> "CacheConfig": + return cls(**d) + + +def hash_quant_config(quant_config: Any) -> str: + """Compute a stable hash of the quantization config. + + Avoids str()/repr() on arbitrary objects because those embed memory + addresses (e.g. "at 0x7f..."), producing different hashes across + processes and causing permanent config mismatch. + """ + if quant_config is None: + return "" + try: + if hasattr(quant_config, "to_dict"): + config_str = json.dumps(quant_config.to_dict(), sort_keys=True) + elif isinstance(quant_config, dict): + config_str = json.dumps(quant_config, sort_keys=True) + elif hasattr(quant_config, "__dict__"): + config_str = ( + type(quant_config).__name__ + + ":" + + json.dumps( + { + k: v + for k, v in sorted(quant_config.__dict__.items()) + if not k.startswith("_") + and isinstance( + v, (str, int, float, bool, type(None), list, dict) + ) + }, + sort_keys=True, + ) + ) + else: + config_str = type(quant_config).__name__ + return hashlib.sha256(config_str.encode()).hexdigest() + except Exception: + config_str = type(quant_config).__name__ + return hashlib.sha256(config_str.encode()).hexdigest() + + +def get_quant_method_name(quant_config: Any) -> str: + """Extract the quantization method name from config.""" + if quant_config is None: + return "" + if isinstance(quant_config, str): + return quant_config + if hasattr(quant_config, "get_name"): + return quant_config.get_name() + if hasattr(quant_config, "name"): + return quant_config.name + return type(quant_config).__name__ + + +# --------------------------------------------------------------------------- +# IPC quantization-method allowlist +# --------------------------------------------------------------------------- +# +# CUDA IPC zero-copy sharing exports ONLY raw tensor data, so it is correct only +# when process_weights_after_loading's entire effect is captured by that data. +# Methods that stamp Python-side metadata (e.g. block-FP8's format_ue8m0) or +# repack/transpose weights into shapes the meta-init client can't reproduce +# (per-tensor FP8, Marlin, AWQ/GPTQ) would serve silently-wrong numerics. Only +# methods verified to round-trip through pure tensor export are allowed; every +# other method hard-errors. Extend the registry below only after verifying a +# method end-to-end. + + +class UnsupportedQuantForIPCError(RuntimeError): + """Raised when a quantization method is not on the verified allowlist for + CUDA IPC zero-copy weight sharing.""" + + +def _get_quant_field(quant_config: Any, key: str) -> Any: + """Read a field from a quant config that may be a dict or an object.""" + if quant_config is None: + return None + if isinstance(quant_config, dict): + return quant_config.get(key) + return getattr(quant_config, key, None) + + +def _fp8_round_trips_via_ipc(quant_config: Any) -> bool: + """Only block-wise FP8 is verified. + + Block-wise FP8 (weight_block_size set) preserves weight shape and the only + post-load metadata it stamps is accounted for. Per-tensor FP8 transposes + `layer.weight` during post-processing, a shape change the meta-init client + cannot reproduce, so it is not supported. + """ + return _get_quant_field(quant_config, "weight_block_size") is not None + + +# quant_method name -> predicate(quant_config) -> bool (True == verified safe). +# A method absent from this registry is unsupported and hard-errors. +IPC_QUANT_ALLOWLIST = { + "": lambda _quant_config: True, # unquantized + "fp8": _fp8_round_trips_via_ipc, # only block-wise FP8 verified +} + + +def is_ipc_quant_supported(quant_method: str, quant_config: Any) -> bool: + """Return True if `quant_method` is verified safe for IPC zero-copy sharing.""" + predicate = IPC_QUANT_ALLOWLIST.get(quant_method) + if predicate is None: + return False + return bool(predicate(quant_config)) + + +def check_ipc_quant_support( + quant_method: str, quant_config: Any, *, where: str +) -> None: + """Hard-error unless `quant_method` is verified safe for IPC zero-copy sharing. + + `where` is a short tag (e.g. "daemon"/"client") used only in the error + message. Raises UnsupportedQuantForIPCError with an actionable message. + """ + if is_ipc_quant_supported(quant_method, quant_config): + return + verified = ", ".join( + (repr(m) if m else "'' (unquantized)") for m in IPC_QUANT_ALLOWLIST + ) + raise UnsupportedQuantForIPCError( + f"[weight_cache:{where}] quantization method {quant_method!r} is not " + f"verified for CUDA IPC zero-copy weight sharing. Its " + f"process_weights_after_loading may stamp Python-side metadata " + f"(e.g. format_ue8m0) or repack/transpose weights into shapes the " + f"meta-initialized client cannot reproduce, which would silently serve " + f"wrong-numerics weights. Verified methods: {verified}. Note: FP8 is " + f"only verified for block-wise configs (weight_block_size set), not " + f"per-tensor FP8. Disable the weight cache (--weight-cache-mode off) " + f"for this model." + ) + + +# --------------------------------------------------------------------------- +# Socket protocol helpers +# --------------------------------------------------------------------------- + + +MAX_MSG_SIZE = 256 * 1024 * 1024 # 256 MiB + + +def send_msg(sock, obj: Any) -> None: + """Send a length-prefixed pickled message over a socket.""" + data = pickle.dumps(obj, protocol=pickle.HIGHEST_PROTOCOL) + header = struct.pack("!I", len(data)) + sock.sendall(header + data) + + +def recv_msg(sock) -> Any: + """Receive a length-prefixed pickled message from a socket.""" + header = _recv_exact(sock, 4) + if header is None: + raise ConnectionError("Connection closed while reading message header") + length = struct.unpack("!I", header)[0] + if length > MAX_MSG_SIZE: + raise ValueError(f"Message size {length} exceeds {MAX_MSG_SIZE} byte cap") + data = _recv_exact(sock, length) + if data is None: + raise ConnectionError("Connection closed while reading message body") + return safe_pickle_loads(data) + + +def _recv_exact(sock, n: int) -> Optional[bytes]: + """Receive exactly n bytes from a socket.""" + buf = bytearray() + while len(buf) < n: + chunk = sock.recv(n - len(buf)) + if not chunk: + return None + buf.extend(chunk) + return bytes(buf) + + +def compute_env_stamp() -> Dict[str, str]: + """Local environment fingerprint for the IPC weight cache. + + Returns the device compute capability and torch version of the current + process. A daemon and a connecting client that differ on either may have run + different post-processing / kernel-selection branches, producing weights that + map cleanly over IPC yet serve garbage; stamping these into CacheConfig turns + that into a clean mismatch. Imported lazily so protocol.py stays cheap to + import and usable on CPU-only hosts (both fields degrade to ""). + """ + device_capability = "" + torch_version = "" + try: + import torch + + torch_version = str(torch.__version__) + except Exception: + pass + try: + from sglang.srt.platforms import current_platform + + cap = current_platform.get_device_capability() + if cap is not None: + device_capability = f"{cap.major}.{cap.minor}" + except Exception: + pass + return {"device_capability": device_capability, "torch_version": torch_version} + + +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. + """ + return tp_size * pp_rank + tp_rank + + +def compute_local_gpu_id( + pp_rank: int, + tp_rank: int, + pp_size_per_node: int, + tp_size_per_node: int, + base_gpu_id: int = 0, + gpu_id_step: int = 1, +) -> int: + """Single source of truth for the local GPU id a daemon rank runs on. + + Mirrors the engine's device assignment so a daemon and the engine rank it + serves always land on the same physical GPU (a prerequisite for CUDA IPC). + ``base_gpu_id``/``gpu_id_step`` default to the identity mapping used by the + standalone launcher; the engine passes its real ``--base-gpu-id`` / + ``--gpu-id-step`` so every call site computes the id the same way instead of + keeping three drifting copies of the formula. + """ + return ( + base_gpu_id + + (pp_rank % pp_size_per_node) * tp_size_per_node + + (tp_rank % tp_size_per_node) * gpu_id_step + ) + + +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 WEIGHT_CACHE_SOCKET_TEMPLATE.format(global_rank=global_rank) + + +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 WEIGHT_CACHE_READY_TEMPLATE.format(global_rank=global_rank) + + +def _read_ready_pid(ready_path: str) -> Optional[int]: + """Read the daemon PID from a .ready file. Returns None if unreadable.""" + try: + with open(ready_path) as f: + for line in f: + if line.startswith("pid="): + return int(line.strip().split("=", 1)[1]) + except (OSError, ValueError): + pass + return None + + +def _is_pid_alive(pid: int) -> bool: + """Check whether a process is still running.""" + try: + os.kill(pid, 0) + return True + except ProcessLookupError: + return False + except PermissionError: + 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. + + 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, + unless ``force`` is set, in which case the running daemon is killed and its + files are taken over (stale-takeover path for a wedged/orphaned daemon). + 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) + + if not os.path.exists(ready_path) and not os.path.exists(socket_path): + return + + pid = _read_ready_pid(ready_path) if os.path.exists(ready_path) else 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"(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." + ) + try: + os.kill(pid, signal.SIGKILL) + except ProcessLookupError: + pass + + for path in (ready_path, socket_path): + if os.path.exists(path): + os.unlink(path) + logger.info(f"Removed stale daemon file: {path}") diff --git a/test/manual/test_weight_cache_e2e.py b/test/manual/test_weight_cache_e2e.py new file mode 100644 index 000000000..562820497 --- /dev/null +++ b/test/manual/test_weight_cache_e2e.py @@ -0,0 +1,377 @@ +"""E2E test for WeightCacheDaemon with real model loading. + +Launches TP daemons that load a real model, export IPC handles, +and verifies the client can fetch and import them. + +Usage: + # With a small model (single GPU): + python test/manual/test_weight_cache_e2e.py --model-path /path/to/model --tp-size 1 + + # With a large model (multi-GPU): + python test/manual/test_weight_cache_e2e.py --model-path /path/to/model --tp-size 4 + +Requires GPUs and the model to be available on disk. +""" + +import argparse +import multiprocessing as mp +import os +import socket +import sys +import time + + +def find_free_port(): + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: + s.bind(("127.0.0.1", 0)) + return s.getsockname()[1] + + +def _global_rank(tp_size, pp_rank, tp_rank): + return tp_size * pp_rank + tp_rank + + +def _temp_path(tp_size, pp_rank, tp_rank, suffix): + return f"/tmp/sglang_weight_cache_rank{_global_rank(tp_size, pp_rank, tp_rank)}{suffix}" + + +def run_single_daemon( + model_path, + gpu_id, + tp_size, + tp_rank, + pp_size, + pp_rank, + dist_init_method, + socket_path, + ready_path, + done_path, + load_format, + dtype, + quantization, + trust_remote_code, +): + """Run a single daemon process for one (pp_rank, tp_rank).""" + import traceback + + from sglang.srt.weight_cache.daemon import WeightCacheDaemon + + daemon = WeightCacheDaemon( + model_path=model_path, + gpu_id=gpu_id, + tp_size=tp_size, + tp_rank=tp_rank, + pp_size=pp_size, + pp_rank=pp_rank, + dp_size=1, + load_format=load_format, + dtype=dtype, + quantization=quantization, + trust_remote_code=trust_remote_code, + dist_init_method=dist_init_method, + ) + daemon.socket_path = socket_path + daemon.ready_path = ready_path + + try: + daemon.load() + except Exception as e: + tb = traceback.format_exc() + print( + f"[Daemon gpu={gpu_id} pp_rank={pp_rank} tp_rank={tp_rank}] " + f"LOAD ERROR: {e}\n{tb}", + flush=True, + ) + with open(ready_path, "w") as f: + f.write(f"ERROR: {e}\n") + return + + with open(ready_path, "w") as f: + f.write(f"pid={os.getpid()}\n") + f.write(f"num_entries={len(daemon.state_entries)}\n") + print( + f"[Daemon gpu={gpu_id} pp_rank={pp_rank} tp_rank={tp_rank}] Ready with " + f"{len(daemon.state_entries)} tensors", + flush=True, + ) + + from sglang.srt.weight_cache.protocol import CacheConfig, recv_msg, send_msg + + if os.path.exists(socket_path): + os.unlink(socket_path) + s = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) + s.bind(socket_path) + s.listen(5) + s.settimeout(60) + + try: + while not os.path.exists(done_path): + try: + conn, _ = s.accept() + try: + req = recv_msg(conn) + if req.get("type") == "query_config": + send_msg( + conn, {"status": "ok", "config": daemon.config.to_dict()} + ) + elif req.get("type") == "fetch_state": + engine_config = CacheConfig.from_dict(req["config"]) + if daemon.config.matches(engine_config): + send_msg( + conn, + { + "status": "ok", + "config": daemon.config.to_dict(), + "entries": daemon.state_entries, + }, + ) + else: + send_msg( + conn, + { + "status": "mismatch", + "daemon_config": daemon.config.to_dict(), + }, + ) + elif req.get("type") == "ping": + send_msg(conn, {"status": "ok"}) + except Exception: + pass + finally: + conn.close() + except socket.timeout: + continue + finally: + s.close() + if os.path.exists(socket_path): + os.unlink(socket_path) + + print( + f"[Daemon gpu={gpu_id} pp_rank={pp_rank} tp_rank={tp_rank}] Exiting", + flush=True, + ) + + +def main(): + parser = argparse.ArgumentParser(description="E2E test for WeightCacheDaemon") + parser.add_argument("--model-path", required=True, help="Path to model weights") + parser.add_argument("--tp-size", type=int, default=1, help="Tensor parallel size") + parser.add_argument("--pp-size", type=int, default=1, help="Pipeline parallel size") + parser.add_argument("--load-format", default="auto", help="Weight load format") + parser.add_argument("--dtype", default="auto", help="Model dtype") + parser.add_argument("--quantization", default=None, help="Quantization method") + parser.add_argument("--trust-remote-code", action="store_true") + parser.add_argument("--timeout", type=int, default=1800, help="Timeout in seconds") + args = parser.parse_args() + + tp_size = args.tp_size + pp_size = args.pp_size + total_ranks = tp_size * pp_size + print(f"=== E2E Test: WeightCacheDaemon TP={tp_size} PP={pp_size} ===") + print(f"Model: {args.model_path}") + + dist_port = find_free_port() + dist_init_method = f"tcp://127.0.0.1:{dist_port}" + print(f"dist_init_method = {dist_init_method}") + + # Clean up old files + for pp_rank in range(pp_size): + for tp_rank in range(tp_size): + for suffix in [".sock", ".ready", ".done"]: + p = _temp_path(tp_size, pp_rank, tp_rank, suffix) + if os.path.exists(p): + os.unlink(p) + + mp.set_start_method("spawn", force=True) + + # Launch all daemon processes + procs = [] + for pp_rank in range(pp_size): + for tp_rank in range(tp_size): + gpu_id = pp_rank * tp_size + tp_rank + sock_path = _temp_path(tp_size, pp_rank, tp_rank, ".sock") + rdy_path = _temp_path(tp_size, pp_rank, tp_rank, ".ready") + done_path = _temp_path(tp_size, pp_rank, tp_rank, ".done") + + p = mp.Process( + target=run_single_daemon, + args=( + args.model_path, + gpu_id, + tp_size, + tp_rank, + pp_size, + pp_rank, + dist_init_method, + sock_path, + rdy_path, + done_path, + args.load_format, + args.dtype, + args.quantization, + args.trust_remote_code, + ), + name=f"daemon_gpu{gpu_id}", + ) + p.start() + procs.append(p) + print( + f"Launched daemon gpu={gpu_id} pp_rank={pp_rank} " + f"tp_rank={tp_rank} pid={p.pid}" + ) + + # Wait for all daemons to be ready + print("Waiting for all daemons to load model...") + start = time.time() + error_found = False + for pp_rank in range(pp_size): + for tp_rank in range(tp_size): + ready_path = _temp_path(tp_size, pp_rank, tp_rank, ".ready") + while not os.path.exists(ready_path): + elapsed = time.time() - start + if elapsed > args.timeout: + print( + f"ERROR: Daemon pp_rank={pp_rank} tp_rank={tp_rank} " + f"timeout after {args.timeout}s" + ) + error_found = True + break + for p in procs: + if not p.is_alive() and not os.path.exists(ready_path): + print( + f"ERROR: Daemon pid={p.pid} exited with code {p.exitcode}" + ) + error_found = True + break + if error_found: + break + time.sleep(2) + + if error_found: + break + + with open(ready_path) as f: + content = f.read() + if content.startswith("ERROR:"): + print( + f"ERROR: Daemon pp_rank={pp_rank} tp_rank={tp_rank} " + f"failed: {content.strip()}" + ) + error_found = True + break + + print( + f"Daemon pp_rank={pp_rank} tp_rank={tp_rank} ready " + f"({time.time()-start:.0f}s)" + ) + if error_found: + break + + if error_found: + for p in procs: + if p.is_alive(): + p.terminate() + sys.exit(1) + + print( + f"\nAll {total_ranks} daemons ready! Total load time: {time.time()-start:.1f}s" + ) + + # Query config from daemon (pp_rank=0, tp_rank=0) + from sglang.srt.utils import MultiprocessingSerializer + from sglang.srt.weight_cache.protocol import recv_msg, send_msg + + socket_path_0 = _temp_path(tp_size, 0, 0, ".sock") + s = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) + s.settimeout(10) + s.connect(socket_path_0) + send_msg(s, {"type": "query_config"}) + result = recv_msg(s) + s.close() + daemon_config = result.get("config", {}) + print( + f"\nDaemon (0,0) config: model={daemon_config.get('model_path')}, " + f"arch={daemon_config.get('model_arch')}, " + f"tp_size={daemon_config.get('tp_size')}, " + f"dtype={daemon_config.get('dtype')}" + ) + + # Fetch state from daemon (pp_rank=0, tp_rank=0) + s = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) + s.settimeout(300) + s.connect(socket_path_0) + send_msg(s, {"type": "fetch_state", "config": daemon_config}) + tic = time.perf_counter() + result = recv_msg(s) + fetch_time = time.perf_counter() - tic + s.close() + + print( + f"\nFetch from daemon (0,0): status={result['status']}, time={fetch_time:.2f}s" + ) + + if result["status"] == "ok": + entries = result["entries"] + print(f"Received {len(entries)} IPC handles from daemon (0,0)") + + sample_names = list(entries.keys())[:5] + for name in sample_names: + entry = entries[name] + imported = MultiprocessingSerializer.deserialize(entry["handle"]) + print( + f" {name}: shape={tuple(imported.shape)}, " + f"dtype={imported.dtype}, device={imported.device}" + ) + del imported + + print("\nIPC import OK!") + else: + print(f"ERROR: fetch_state returned {result}") + for pp_rank in range(pp_size): + for tp_rank in range(tp_size): + done = _temp_path(tp_size, pp_rank, tp_rank, ".done") + with open(done, "w") as f: + f.write("done\n") + for p in procs: + if p.is_alive(): + p.terminate() + sys.exit(1) + + # Test config mismatch + mismatch_config = dict(daemon_config) + mismatch_config["tp_size"] = daemon_config.get("tp_size", 1) + 1 + s = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) + s.settimeout(10) + s.connect(socket_path_0) + send_msg(s, {"type": "fetch_state", "config": mismatch_config}) + result = recv_msg(s) + s.close() + if result["status"] == "mismatch": + print("Config mismatch detection: OK") + else: + print(f"WARNING: Expected mismatch, got {result['status']}") + + # Signal all daemons done + for pp_rank in range(pp_size): + for tp_rank in range(tp_size): + done = _temp_path(tp_size, pp_rank, tp_rank, ".done") + with open(done, "w") as f: + f.write("done\n") + for p in procs: + p.join(timeout=15) + if p.is_alive(): + p.terminate() + + # Clean up + for pp_rank in range(pp_size): + for tp_rank in range(tp_size): + for suffix in [".sock", ".ready", ".done"]: + path = _temp_path(tp_size, pp_rank, tp_rank, suffix) + if os.path.exists(path): + os.unlink(path) + + print("\n=== E2E Test Passed! ===") + + +if __name__ == "__main__": + main() diff --git a/test/registered/model_loading/test_weight_cache_daemon.py b/test/registered/model_loading/test_weight_cache_daemon.py new file mode 100644 index 000000000..a49dfc446 --- /dev/null +++ b/test/registered/model_loading/test_weight_cache_daemon.py @@ -0,0 +1,338 @@ +import os +import subprocess +import sys +import time +import unittest + +import requests +import torch + +from sglang.srt.utils import kill_process_tree +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import ( + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + DEFAULT_URL_FOR_TEST, + CustomTestCase, + popen_launch_server, +) + +# A ~1B model keeps the daemon->client IPC handoff cheap to exercise on every +# PR (fast download + load) while still covering the real block-load path; the +# test asserts the IPC path ran, not any particular model's quality. +DEFAULT_MODEL = "Qwen/Qwen3-0.6B" + +# This file runs in two suites. The TP=2 class needs the 2-GPU runner (extra-a); +# the TP=1 smoke class is always-on (base-b / 1-gpu-small) so the daemon->client +# IPC handoff is exercised on every PR. Since the CI runner executes the whole +# file per suite, TestWeightCacheDaemonTP2 self-skips when fewer than 2 GPUs are +# visible (i.e. on the 1-gpu runner). +register_cuda_ci(est_time=100, stage="extra-a", runner_config="2-gpu-large") +register_cuda_ci(est_time=100, stage="base-b", runner_config="1-gpu-small") + +# Capture the client server's logs so test_loaded_via_ipc can assert the IPC +# load path actually ran (and did not silently fall back to disk). +STDOUT_FILENAME = "/tmp/test_weight_cache_daemon_stdout.log" +STDERR_FILENAME = "/tmp/test_weight_cache_daemon_stderr.log" +SMOKE_STDOUT_FILENAME = "/tmp/test_weight_cache_daemon_smoke_stdout.log" +SMOKE_STDERR_FILENAME = "/tmp/test_weight_cache_daemon_smoke_stderr.log" + +PROMPTS = [ + "The capital of France is", + "Hello, my name is", + "The future of AI is", +] + + +@unittest.skipIf( + torch.cuda.device_count() < 2, + "TP=2 weight cache daemon test requires >=2 GPUs (skipped on the 1-gpu runner)", +) +class TestWeightCacheDaemonTP2(CustomTestCase): + """E2E test: start weight cache daemons, then launch server in client mode with TP2.""" + + @classmethod + def setUpClass(cls): + cls.model = DEFAULT_MODEL + cls.base_url = DEFAULT_URL_FOR_TEST + cls.tp_size = 2 + + # 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}" + if os.path.exists(path): + os.unlink(path) + + # Step 1: Launch weight cache daemons (blocks until all ranks are ready, + # then monitors child processes) + cls.daemon_process = subprocess.Popen( + [ + sys.executable, + "-m", + "sglang.srt.weight_cache.daemon", + "--model-path", + cls.model, + "--tp-size", + str(cls.tp_size), + ] + ) + + # 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" + 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" + ) + if cls.daemon_process.poll() is not None: + raise RuntimeError( + f"Weight cache daemon exited prematurely " + f"with code {cls.daemon_process.returncode}" + ) + time.sleep(2) + + # Step 3: Launch server in client mode — loads weights via IPC from daemons + cls.stdout = open(STDOUT_FILENAME, "w") + cls.stderr = open(STDERR_FILENAME, "w") + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=[ + "--tp", + str(cls.tp_size), + "--weight-cache-mode", + "client", + ], + return_stdout_stderr=(cls.stdout, cls.stderr), + ) + + @classmethod + def tearDownClass(cls): + if hasattr(cls, "process") and cls.process: + kill_process_tree(cls.process.pid) + if hasattr(cls, "daemon_process") and cls.daemon_process: + kill_process_tree(cls.daemon_process.pid) + for stream in (getattr(cls, "stdout", None), getattr(cls, "stderr", None)): + if stream is not None: + try: + stream.close() + except OSError: + pass + for path in (STDOUT_FILENAME, STDERR_FILENAME): + if os.path.exists(path): + try: + 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}" + if os.path.exists(path): + try: + os.unlink(path) + except OSError: + pass + + def test_generate(self): + for prompt in PROMPTS: + resp = requests.post( + f"{self.base_url}/v1/completions", + json={ + "model": self.model, + "prompt": prompt, + "max_tokens": 32, + "temperature": 0, + }, + ) + self.assertEqual(resp.status_code, 200) + data = resp.json() + text = data["choices"][0]["text"] + self.assertIsInstance(text, str) + self.assertGreater(len(text), 0, f"Empty output for prompt: {prompt}") + + def test_chat(self): + resp = requests.post( + f"{self.base_url}/v1/chat/completions", + json={ + "model": self.model, + "messages": [{"role": "user", "content": "What is 2+3?"}], + "max_tokens": 32, + "temperature": 0, + }, + ) + self.assertEqual(resp.status_code, 200) + data = resp.json() + content = data["choices"][0]["message"]["content"] + self.assertIsInstance(content, str) + self.assertGreater(len(content), 0) + + def test_loaded_via_ipc(self): + """Assert the server actually loaded weights over IPC. + + Without this, the test would still pass if the IPC path silently + regressed to disk loading (the daemon would just sit unused), because + generation output looks identical either way. The daemon-side loader + logs "[IpcModelLoader] Loaded model via IPC" on every rank, so its + presence in the captured server logs is our proof the IPC path ran. + """ + for stream in (getattr(self, "stdout", None), getattr(self, "stderr", None)): + if stream is not None: + try: + stream.flush() + except OSError: + pass + logs = "" + for path in (STDOUT_FILENAME, STDERR_FILENAME): + if os.path.exists(path): + with open(path, errors="replace") as f: + logs += f.read() + self.assertIn( + "Loaded model via IPC", + logs, + "Expected the client server to load weights via IPC, but the IPC " + "load log line was not found — the loader likely fell back to disk.", + ) + + +class TestWeightCacheDaemonTP1Smoke(CustomTestCase): + """Always-on TP=1 smoke: start a single weight cache daemon, launch a server + in client mode, and confirm it loads weights via IPC and generates. + + This is the fast single-GPU sanity check (small model) that runs on every PR + in the base-b / 1-gpu-small suite; the heavier TP=2 case above only runs on + the 2-GPU runner. + """ + + @classmethod + def setUpClass(cls): + cls.model = DEFAULT_MODEL + cls.base_url = DEFAULT_URL_FOR_TEST + cls.tp_size = 1 + + # 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}" + if os.path.exists(path): + os.unlink(path) + + # Step 1: Launch the weight cache daemon (blocks until the rank is + # ready, then monitors the child process). + cls.daemon_process = subprocess.Popen( + [ + sys.executable, + "-m", + "sglang.srt.weight_cache.daemon", + "--model-path", + cls.model, + "--tp-size", + str(cls.tp_size), + ] + ) + + # 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" + 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" + ) + if cls.daemon_process.poll() is not None: + raise RuntimeError( + f"Weight cache daemon exited prematurely " + f"with code {cls.daemon_process.returncode}" + ) + time.sleep(2) + + # Step 3: Launch server in client mode — loads weights via IPC. + cls.stdout = open(SMOKE_STDOUT_FILENAME, "w") + cls.stderr = open(SMOKE_STDERR_FILENAME, "w") + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=[ + "--tp", + str(cls.tp_size), + "--weight-cache-mode", + "client", + ], + return_stdout_stderr=(cls.stdout, cls.stderr), + ) + + @classmethod + def tearDownClass(cls): + if hasattr(cls, "process") and cls.process: + kill_process_tree(cls.process.pid) + if hasattr(cls, "daemon_process") and cls.daemon_process: + kill_process_tree(cls.daemon_process.pid) + for stream in (getattr(cls, "stdout", None), getattr(cls, "stderr", None)): + if stream is not None: + try: + stream.close() + except OSError: + pass + for path in (SMOKE_STDOUT_FILENAME, SMOKE_STDERR_FILENAME): + if os.path.exists(path): + try: + 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}" + if os.path.exists(path): + try: + os.unlink(path) + except OSError: + pass + + def test_generate(self): + resp = requests.post( + f"{self.base_url}/v1/completions", + json={ + "model": self.model, + "prompt": "The capital of France is", + "max_tokens": 32, + "temperature": 0, + }, + ) + self.assertEqual(resp.status_code, 200) + data = resp.json() + text = data["choices"][0]["text"] + self.assertIsInstance(text, str) + self.assertGreater(len(text), 0, "Empty generation output") + + def test_loaded_via_ipc(self): + """Assert the server actually loaded weights over IPC (see the TP=2 + variant for why this guard matters).""" + for stream in (getattr(self, "stdout", None), getattr(self, "stderr", None)): + if stream is not None: + try: + stream.flush() + except OSError: + pass + logs = "" + for path in (SMOKE_STDOUT_FILENAME, SMOKE_STDERR_FILENAME): + if os.path.exists(path): + with open(path, errors="replace") as f: + logs += f.read() + self.assertIn( + "Loaded model via IPC", + logs, + "Expected the client server to load weights via IPC, but the IPC " + "load log line was not found — the loader likely fell back to disk.", + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/model_loader/test_weight_cache_protocol.py b/test/registered/unit/model_loader/test_weight_cache_protocol.py new file mode 100644 index 000000000..4788897fd --- /dev/null +++ b/test/registered/unit/model_loader/test_weight_cache_protocol.py @@ -0,0 +1,366 @@ +""" +CPU-only unit tests for the weight cache protocol layer. + +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 + - quant-config hashing and method-name extraction + - the daemon rank formula and socket/ready path derivation + - the IPC quantization allowlist (the gate that keeps silently-wrong + quant methods off the zero-copy path) + - stale-vs-live daemon file cleanup + +They intentionally require no CUDA, no model download, and no daemon +process, so they run in the fast CPU suite and would catch a regression +in any of these branches before it reaches the expensive GPU path. +""" + +import os +import socket +import struct +import unittest + +from sglang.srt.weight_cache.protocol import ( + IPC_QUANT_ALLOWLIST, + CacheConfig, + UnsupportedQuantForIPCError, + check_ipc_quant_support, + cleanup_stale_daemon_files, + compute_global_rank, + compute_local_gpu_id, + get_quant_method_name, + get_ready_path, + get_socket_path, + hash_quant_config, + is_ipc_quant_supported, + recv_msg, + send_msg, +) +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=10, suite="base-a-test-cpu") + + +def _make_cache_config(**overrides) -> CacheConfig: + base = dict( + model_path="/models/demo", + model_arch="LlamaForCausalLM", + tp_size=2, + tp_rank=0, + pp_size=1, + pp_rank=0, + dp_size=1, + ep_size=1, + quant_method="", + quant_config_hash="", + dtype="torch.float16", + revision="", + device_capability="8.0", + torch_version="2.5.1", + ) + base.update(overrides) + return CacheConfig(**base) + + +class TestProtocolFraming(CustomTestCase): + """Length-prefixed pickle framing over a real socket pair.""" + + def test_round_trip(self): + a, b = socket.socketpair() + try: + payload = {"handles": [1, 2, 3], "meta": ("x", 4.5), "flag": True} + send_msg(a, payload) + self.assertEqual(recv_msg(b), payload) + finally: + a.close() + b.close() + + def test_multiple_messages_are_framed_independently(self): + a, b = socket.socketpair() + try: + send_msg(a, {"n": 1}) + send_msg(a, {"n": 2}) + self.assertEqual(recv_msg(b), {"n": 1}) + self.assertEqual(recv_msg(b), {"n": 2}) + finally: + a.close() + b.close() + + def test_connection_closed_mid_header_raises(self): + a, b = socket.socketpair() + try: + # Peer sends only a partial header, then hangs up. + a.sendall(struct.pack("!I", 128)[:2]) + a.close() + with self.assertRaises(ConnectionError): + recv_msg(b) + finally: + b.close() + + def test_connection_closed_mid_body_raises(self): + a, b = socket.socketpair() + try: + # Full header promising 128 bytes, but no body follows. + a.sendall(struct.pack("!I", 128)) + a.close() + with self.assertRaises(ConnectionError): + recv_msg(b) + finally: + b.close() + + +class TestCacheConfig(CustomTestCase): + def test_identical_configs_match(self): + self.assertTrue(_make_cache_config().matches(_make_cache_config())) + + def test_any_field_difference_breaks_match(self): + base = _make_cache_config() + for field, value in ( + ("tp_rank", 1), + ("dtype", "torch.bfloat16"), + ("quant_method", "fp8"), + ("model_path", "/models/other"), + ("revision", "v2"), + ("device_capability", "9.0"), + ("torch_version", "2.4.0"), + ): + self.assertFalse( + base.matches(_make_cache_config(**{field: value})), + msg=f"{field} difference should break match", + ) + + def test_dict_round_trip(self): + cfg = _make_cache_config(quant_method="fp8", quant_config_hash="abc123") + restored = CacheConfig.from_dict(cfg.to_dict()) + self.assertTrue(cfg.matches(restored)) + self.assertEqual(cfg.to_dict(), restored.to_dict()) + + +class TestQuantConfigHashing(CustomTestCase): + def test_none_hashes_to_empty(self): + self.assertEqual(hash_quant_config(None), "") + + def test_dict_hash_is_deterministic_and_order_insensitive(self): + h1 = hash_quant_config({"bits": 8, "group_size": 128}) + h2 = hash_quant_config({"group_size": 128, "bits": 8}) + self.assertEqual(h1, h2) + self.assertNotEqual(h1, hash_quant_config({"bits": 4, "group_size": 128})) + + def test_hash_is_not_truncated(self): + # A correctness gate must use the full SHA-256 digest, not a 16-char prefix. + self.assertEqual(len(hash_quant_config({"bits": 8})), 64) + + def test_hash_does_not_embed_object_address(self): + # Two distinct instances with identical public attrs must hash equal, + # otherwise configs would mismatch across processes (the bug the + # docstring warns about). + class _Q: + def __init__(self): + self.bits = 8 + self.method = "fp8" + + self.assertEqual(hash_quant_config(_Q()), hash_quant_config(_Q())) + + def test_get_quant_method_name_variants(self): + self.assertEqual(get_quant_method_name(None), "") + self.assertEqual(get_quant_method_name("fp8"), "fp8") + + class _WithGetName: + def get_name(self): + return "gptq_marlin" + + class _WithName: + name = "awq" + + self.assertEqual(get_quant_method_name(_WithGetName()), "gptq_marlin") + self.assertEqual(get_quant_method_name(_WithName()), "awq") + + +class TestGlobalRankAndPaths(CustomTestCase): + def test_compute_global_rank_formula(self): + self.assertEqual(compute_global_rank(tp_size=4, pp_rank=0, tp_rank=3), 3) + 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_compute_local_gpu_id_honors_base_and_step(self): + # Single-node TP=4: identity mapping rank -> gpu. + self.assertEqual( + compute_local_gpu_id(0, 2, pp_size_per_node=1, tp_size_per_node=4), + 2, + ) + # base_gpu_id offsets every rank; gpu_id_step strides between them. + self.assertEqual( + compute_local_gpu_id( + 0, 2, pp_size_per_node=1, tp_size_per_node=4, base_gpu_id=4 + ), + 6, + ) + self.assertEqual( + compute_local_gpu_id( + 0, 2, pp_size_per_node=1, tp_size_per_node=4, gpu_id_step=2 + ), + 4, + ) + + +class TestIpcQuantAllowlist(CustomTestCase): + def test_unquantized_is_supported(self): + self.assertTrue(is_ipc_quant_supported("", None)) + + def test_block_fp8_supported_but_per_tensor_fp8_rejected(self): + self.assertTrue( + is_ipc_quant_supported("fp8", {"weight_block_size": [128, 128]}) + ) + # Per-tensor FP8 (no weight_block_size) transposes the weight during + # post-processing -> not reproducible by the meta-init client. + self.assertFalse(is_ipc_quant_supported("fp8", {})) + self.assertFalse(is_ipc_quant_supported("fp8", None)) + + def test_unknown_method_rejected(self): + self.assertFalse(is_ipc_quant_supported("gptq_marlin", None)) + self.assertFalse(is_ipc_quant_supported("awq", None)) + + def test_check_raises_on_unsupported(self): + with self.assertRaises(UnsupportedQuantForIPCError): + check_ipc_quant_support("awq", None, where="client") + # Per-tensor FP8 must also raise even though "fp8" is a known key. + with self.assertRaises(UnsupportedQuantForIPCError): + check_ipc_quant_support("fp8", {}, where="daemon") + + def test_check_passes_on_supported(self): + # Should not raise. + check_ipc_quant_support("", None, where="daemon") + check_ipc_quant_support( + "fp8", {"weight_block_size": [128, 128]}, where="daemon" + ) + + def test_allowlist_registry_shape(self): + # Guard against accidentally widening the allowlist without review. + self.assertEqual(set(IPC_QUANT_ALLOWLIST), {"", "fp8"}) + + +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 + + def _paths(self): + return get_ready_path(self.RANK), get_socket_path(self.RANK) + + def tearDown(self): + for path in self._paths(): + if os.path.exists(path): + os.unlink(path) + + def test_no_files_is_noop(self): + # Neither file present: must return quietly, not raise. + cleanup_stale_daemon_files(self.RANK) + + def test_stale_files_without_live_pid_are_removed(self): + ready_path, socket_path = self._paths() + # A .ready file whose PID is unreadable is treated as a crashed-daemon + # leftover and cleaned up. + with open(ready_path, "w") as f: + f.write("stale contents, no pid line\n") + open(socket_path, "w").close() + + cleanup_stale_daemon_files(self.RANK) + + self.assertFalse(os.path.exists(ready_path)) + self.assertFalse(os.path.exists(socket_path)) + + def test_live_daemon_pid_blocks_cleanup(self): + ready_path, socket_path = self._paths() + # Our own PID is alive -> cleanup must refuse and leave files intact. + with open(ready_path, "w") as f: + f.write(f"pid={os.getpid()}\n") + open(socket_path, "w").close() + + with self.assertRaises(RuntimeError): + cleanup_stale_daemon_files(self.RANK) + + self.assertTrue(os.path.exists(ready_path)) + self.assertTrue(os.path.exists(socket_path)) + + def test_force_takes_over_from_live_pid(self): + ready_path, socket_path = self._paths() + # Spawn a real child we are allowed to kill, point the ready file at it, + # then force-takeover: the child must be killed and the files removed. + import subprocess + import sys + + child = subprocess.Popen([sys.executable, "-c", "import time; time.sleep(30)"]) + try: + with open(ready_path, "w") as f: + f.write(f"pid={child.pid}\n") + open(socket_path, "w").close() + + cleanup_stale_daemon_files(self.RANK, 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. + self.assertEqual(child.wait(timeout=5), -9) + finally: + if child.poll() is None: + child.kill() + child.wait(timeout=5) + + +class TestDaemonModeRefusesDiskLoad(CustomTestCase): + """In daemon mode the engine and daemon share a GPU, so a missing daemon + must be a hard error — NOT a silent disk-load that would OOM the shared + device. This exercises that contract without a GPU or a live daemon by + pointing the loader at a socket path that does not exist. + """ + + RANK = 987655 + + def _model_config(self): + from types import SimpleNamespace + + # Minimal stand-in: the loader only reads hf_config.quantization_config, + # quantization, and (unreached here) hf_config.architectures. + hf_config = SimpleNamespace( + architectures=["LlamaForCausalLM"], quantization_config=None + ) + return SimpleNamespace( + model_path="/models/demo", + hf_config=hf_config, + quantization=None, + revision=None, + dtype="torch.float16", + ) + + def test_daemon_mode_missing_daemon_raises_instead_of_disk_load(self): + 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) + if os.path.exists(missing_socket): + os.unlink(missing_socket) + + loader = IpcModelLoader( + load_config=LoadConfig(load_format=LoadFormat.IPC_CACHE), + socket_path=missing_socket, + weight_cache_mode="daemon", + fallback_load_format="auto", + ) + + with self.assertRaises(RuntimeError) as ctx: + loader.load_model(model_config=self._model_config(), device_config=None) + # The error must be about the missing daemon, proving we did not quietly + # fall through to a disk load. + self.assertIn("daemon", str(ctx.exception).lower()) + + +if __name__ == "__main__": + unittest.main()