class IpcModelLoader(BaseModelLoader):
"""Loads a model by mapping the weight cache daemon's tensors via CUDA IPC.
The model is initialized on the meta device and every parameter/buffer is
replaced by the daemon's post-quantized tensor, so
process_weights_after_loading is skipped entirely. In "zero_copy" mode the
engine shares the daemon's GPU memory; in "copy" mode the tensors are
cloned into engine-owned memory and the daemon is asked to release its
cache afterwards.
Extra config keys (via --model-loader-extra-config):
- socket_path: explicit daemon socket path. Defaults to a per-GPU path
derived from the physical GPU uuid and the cache role (target/draft).
- socket_dir: directory containing the daemon sockets.
- mode: "zero_copy" (default) or "copy".
- fallback: fall back to disk loading when the daemon is unavailable or
the fingerprints mismatch (default: True).
- connect_timeout_s: socket connect timeout (default: 5.0).
- state_timeout_s: timeout for the weight-transfer request (default: 300.0).
Note: in zero-copy mode the weights live in the daemon's CUDA IPC
allocations, so sleep mode (CuMemAllocator weight offloading) must not be
used with this loader.
"""
def __init__(self, load_config: LoadConfig):
super().__init__(load_config)
extra_config = copy(load_config.model_loader_extra_config or {})
self.socket_path: str | None = extra_config.pop("socket_path", None)
self.socket_dir: str | None = extra_config.pop("socket_dir", None)
# Internal: set by the engine when routing a speculative draft to the
# daemon's draft group.
self.is_draft = bool(extra_config.pop("is_draft", False))
if self.is_draft and self.socket_path is not None:
raise ValueError(
"socket_path cannot be combined with the draft weight cache role; "
"use socket_dir so the target and draft sockets are derived "
"independently"
)
self.mode: str = extra_config.pop("mode", "zero_copy")
self.fallback: bool = extra_config.pop("fallback", True)
self.connect_timeout_s: float = float(
extra_config.pop("connect_timeout_s", _CONNECT_TIMEOUT_S)
)
self.state_timeout_s: float = float(
extra_config.pop("state_timeout_s", _STATE_TIMEOUT_S)
)
if self.mode not in ("zero_copy", "copy"):
raise ValueError(
f"Invalid weight cache mode {self.mode!r}, "
"expected 'zero_copy' or 'copy'"
)
if extra_config:
raise ValueError(
f"Unexpected extra config keys for load format "
f"{load_config.load_format}: {sorted(extra_config)}"
)
def get_external_weight_memory(self, vllm_config: VllmConfig) -> int:
# Copy mode clones the weights into this process; nothing external.
if self.mode != "zero_copy":
return 0
total = self._daemon_memory()
if is_draft_model_cacheable(vllm_config.speculative_config):
# The draft group is queried with the same config the engine
# would load the draft with.
draft_config = get_draft_load_config(vllm_config)
if draft_config.load_format == "ipc_cache":
draft_loader = IpcModelLoader(draft_config)
# Only a zero-copy draft-group daemon holds external weights;
# an explicit draft_load_config without is_draft would resolve
# back to the target socket and double-count it.
if draft_loader.is_draft and draft_loader.mode == "zero_copy":
total += draft_loader._daemon_memory()
return total
def _daemon_memory(self) -> int:
if not self.fallback:
# The loader waits for the daemon at load time, so the weights
# will be zero-copy mapped for sure; mirror that wait here since
# returning 0 would over-grant the memory budget.
return self._with_startup_wait(self._query_daemon_memory)
# An unreachable daemon means a disk load, i.e. nothing external.
try:
return self._query_daemon_memory()
except (WeightCacheUnavailableError, ConnectionError, OSError) as e:
logger.warning(
"Cannot query weight cache daemon memory (%s); "
"assuming the weights are not externally held",
e,
)
return 0
def _query_daemon_memory(self) -> int:
with self._connect(self.connect_timeout_s) as conn:
send_msg(conn, {"cmd": "get_memory"})
response = recv_msg(conn)
if response.get("status") != "ok":
raise WeightCacheUnavailableError(
"Weight cache daemon rejected the memory query: "
f"{response.get('message')}"
)
return int(response.get("memory_bytes", 0))
def download_model(self, model_config: ModelConfig) -> None:
DefaultModelLoader(self._fallback_load_config()).download_model(model_config)
def load_weights(self, model: nn.Module, model_config: ModelConfig) -> None:
"""Best-effort in-place reload for an already-initialized model.
Copies daemon tensors into matching parameters/buffers. The model is
expected to already be in the post-quantized layout (e.g. previously
loaded through this loader).
"""
device_index = torch.accelerator.current_device_index()
entries = self._fetch_entries(model_config).entries
params = dict(model.named_parameters())
buffers = dict(model.named_buffers())
for name, entry in entries.items():
target = params.get(name, buffers.get(name))
source = entry.rebuild(device_index)
if target is None or target.shape != source.shape:
logger.warning("Skipping mismatched cached tensor %s", name)
continue
target.data.copy_(source)
@instrument(span_name="Load model")
def load_model(
self, vllm_config: VllmConfig, model_config: ModelConfig, prefix: str = ""
) -> nn.Module:
# An unsupported platform is a permanent misconfiguration rather than
# a transient daemon outage, so it is raised even when fallback is on.
check_ipc_platform_support()
state_fetched = False
try:
# Cross-check the routing flag against the identity of the model
# being loaded: a draft load that lost its flag (or a target load
# that got one) would hit the wrong daemon group and
# fingerprint-mismatch.
spec = vllm_config.speculative_config
inferred = spec is not None and model_config is spec.draft_model_config
if inferred != self.is_draft:
raise CacheConfigMismatchError(
f"Weight cache role mismatch: loading "
f"{'draft' if inferred else 'target'} model but the loader "
f"was configured for the "
f"{'draft' if self.is_draft else 'target'} group"
)
state = self._fetch_entries(model_config)
state_fetched = True
return self._build_model(vllm_config, model_config, prefix, state)
except (WeightCacheUnavailableError, CacheConfigMismatchError) as e:
if not self.fallback:
raise
logger.warning(
"Weight cache unusable (%s); falling back to disk loading", e
)
except UnsupportedQuantForIPCError:
# Unsupported quantization is a permanent misconfiguration rather
# than a transient daemon outage, so it is raised even when
# fallback is on.
raise
except Exception:
if not self.fallback:
raise
logger.exception(
"Weight cache IPC loading failed; falling back to disk loading"
)
# _build_model failed after fetching state without reaching its
# copy-mode release, so the daemon still holds the full cache;
# release it so the disk fallback does not OOM against it.
if state_fetched and self.mode == "copy":
self._send_release()
torch.accelerator.empty_cache()
return self._fallback_load(vllm_config, model_config, prefix)
def _build_model(
self,
vllm_config: VllmConfig,
model_config: ModelConfig,
prefix: str,
state: WeightCacheState,
) -> nn.Module:
device_config = vllm_config.device_config
load_device = (
device_config.device
if self.load_config.device is None
else self.load_config.device
)
target_device = torch.device(load_device)
device_index = (
target_device.index
if target_device.index is not None
else torch.accelerator.current_device_index()
)
with set_default_torch_dtype(model_config.dtype):
with torch.device("meta"):
model = initialize_model(
vllm_config=vllm_config,
model_config=model_config,
prefix=prefix,
)
check_ipc_quant_support(model)
self._apply_entries(model, state, device_index)
# Flags that load_weights would have set (e.g. EAGLE ownership of
# embed_tokens / lm_head); the daemon ran it, this process did not.
for name, value in state.attrs.items():
setattr(model, name, value)
# The daemon exports tensors that already went through
# process_weights_after_loading; re-run it in pre-processed mode
# so quant methods only rebuild Python-side state (e.g. the MoE
# kernel). Leftovers are materialized afterwards so that
# placeholders the daemon-side post-processing consumed are
# dropped rather than filled with uninitialized memory.
with weights_already_processed():
process_weights_after_loading(model, model_config, target_device)
_materialize_remaining_meta_tensors(
model, torch.device(target_device.type, device_index)
)
if self.mode == "copy":
self._send_release()
logger.info(
"Mapped %d tensors from the weight cache daemon (%s mode)",
len(state.entries),
self.mode,
)
return model.eval()
def _apply_entries(
self,
model: nn.Module,
state: WeightCacheState,
device_index: int,
) -> None:
# remove_duplicate=False keeps tied module aliases reachable by name:
# a tied lm_head *is* the embedding module, so the deduplicated view
# would not contain "lm_head" at all.
modules = dict(model.named_modules(remove_duplicate=False))
registered: dict[str, torch.Tensor] = {}
def _register(name: str, tensor: torch.Tensor, is_param: bool) -> None:
module_name, _, leaf = name.rpartition(".")
module = modules.get(module_name)
if module is None:
raise RuntimeError(f"Cached tensor {name} has no matching module")
# Replace via registration rather than param.data assignment,
# which fails for meta tensors. Entries may also introduce
# post-quantization tensors absent from the meta model.
module._parameters.pop(leaf, None)
module._buffers.pop(leaf, None)
if is_param:
obj: torch.Tensor = (
tensor
if isinstance(tensor, nn.Parameter)
else nn.Parameter(tensor, requires_grad=False)
)
module.register_parameter(leaf, obj)
else:
obj = tensor
module.register_buffer(leaf, obj)
registered[name] = obj
for name, entry in state.entries.items():
tensor = entry.rebuild(device_index)
if self.mode == "copy":
tensor = tensor.clone()
_register(name, tensor, entry.kind == "param")
# Re-establish tied-weight aliases by registering the *same* object the
# canonical name resolved to, so parameter identity (and the tie) is
# preserved instead of allocating uninitialized memory.
for alias_name, canonical_name in state.aliases.items():
obj = registered.get(canonical_name)
if obj is None:
logger.warning(
"Cached alias %s references missing canonical tensor %s",
alias_name,
canonical_name,
)
continue
_register(alias_name, obj, isinstance(obj, nn.Parameter))
def _fetch_entries(self, model_config: ModelConfig) -> WeightCacheState:
dp_group = get_dp_group()
pp_group = get_pp_group()
cache_config = WeightCacheKey.from_model_config(
model_config,
tp_size=get_tensor_model_parallel_world_size(),
tp_rank=get_tensor_model_parallel_rank(),
pp_size=pp_group.world_size,
pp_rank=pp_group.rank_in_group,
dp_size=dp_group.world_size,
dp_rank=dp_group.rank_in_group,
is_draft=self.is_draft,
)
if not self.fallback:
return self._request_state_with_startup_wait(cache_config)
return self._request_state(cache_config)
def _request_state_with_startup_wait(
self, cache_config: WeightCacheKey
) -> WeightCacheState:
return self._with_startup_wait(lambda: self._request_state(cache_config))
def _with_startup_wait(self, op: Callable[[], _T]) -> _T:
"""Retry op until the daemon answers or the state timeout elapses;
the daemon may still be loading the model when the engine starts."""
deadline = time.monotonic() + self.state_timeout_s
while True:
try:
return op()
except (WeightCacheUnavailableError, ConnectionError, OSError) as e:
if time.monotonic() >= deadline:
raise WeightCacheUnavailableError(
"Weight cache daemon did not become ready within "
f"{self.state_timeout_s:.1f}s: {e}"
) from e
logger.info_once(
"Waiting up to %.1fs for the weight cache daemon to start",
self.state_timeout_s,
)
time.sleep(
max(
0.0,
min(_STARTUP_RETRY_INTERVAL_S, deadline - time.monotonic()),
)
)
def _request_state(self, cache_config: WeightCacheKey) -> WeightCacheState:
with self._connect(self.state_timeout_s) as conn:
send_msg(conn, {"cmd": "get_state", "cache_config": cache_config})
response = recv_msg(conn)
status = response.get("status")
if status == "mismatch":
raise CacheConfigMismatchError(
f"WeightCacheKey mismatch on fields: {response.get('fields')}"
)
if status != "ok":
raise WeightCacheUnavailableError(
f"Weight cache daemon error: {response.get('message')}"
)
self._check_gpu_uuid(response.get("gpu_uuid"))
return WeightCacheState(
entries=response["entries"],
aliases=response.get("aliases", {}),
attrs=response.get("attrs", {}),
)
def _connect(self, timeout: float) -> socket.socket:
socket_path = self._resolve_socket_path()
# The auto-derived per-user directory is locked to 0700 and checked
# strictly. When the operator explicitly configures a path they own the
# trust decision, so only ownership/symlink safety is enforced.
strict_perms = self.socket_path is None and self.socket_dir is None
try:
verify_socket_owner(socket_path, strict_perms=strict_perms)
except OSError as e:
raise WeightCacheUnavailableError(
f"Weight cache socket {socket_path} is unavailable: {e}"
) from e
sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
sock.settimeout(timeout)
try:
sock.connect(socket_path)
except OSError as e:
sock.close()
raise WeightCacheUnavailableError(
f"Cannot connect to weight cache daemon at {socket_path}: {e}"
) from e
return sock
def _resolve_socket_path(self) -> str:
if self.socket_path is not None:
return self.socket_path
return get_socket_path(
get_current_device_uuid(),
self.socket_dir,
is_draft=self.is_draft,
)
def _check_gpu_uuid(self, daemon_uuid: str | None) -> None:
if daemon_uuid is None:
return
local_uuid = get_current_device_uuid()
if daemon_uuid != local_uuid:
raise CacheConfigMismatchError(
f"Daemon GPU {daemon_uuid} != engine GPU {local_uuid}; "
"check the socket path / GPU mapping"
)
def _send_release(self) -> None:
try:
with self._connect(self.connect_timeout_s) as conn:
send_msg(conn, {"cmd": "release"})
recv_msg(conn)
except (WeightCacheUnavailableError, ConnectionError, OSError):
logger.warning("Failed to ask the weight cache daemon to release")
def _fallback_load_config(self) -> LoadConfig:
# DefaultModelLoader must not see load_format="ipc_cache" or the ipc
# extra config keys.
return dataclasses.replace(
self.load_config,
load_format="auto",
model_loader_extra_config={},
)
def _fallback_load(
self, vllm_config: VllmConfig, model_config: ModelConfig, prefix: str
) -> nn.Module:
loader = DefaultModelLoader(self._fallback_load_config())
return loader.load_model(
vllm_config=vllm_config, model_config=model_config, prefix=prefix
)