autoinference-utils 0.2.9__tar.gz → 0.2.11__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (18) hide show
  1. {autoinference_utils-0.2.9 → autoinference_utils-0.2.11}/PKG-INFO +1 -1
  2. {autoinference_utils-0.2.9 → autoinference_utils-0.2.11}/pyproject.toml +1 -1
  3. autoinference_utils-0.2.11/src/autoinference_utils/cluster/__init__.py +1 -0
  4. autoinference_utils-0.2.11/src/autoinference_utils/cluster/_net.py +19 -0
  5. autoinference_utils-0.2.11/src/autoinference_utils/cluster/config.py +11 -0
  6. autoinference_utils-0.2.11/src/autoinference_utils/cluster/image.py +126 -0
  7. autoinference_utils-0.2.11/src/autoinference_utils/cluster/modelexpress/__init__.py +18 -0
  8. autoinference_utils-0.2.11/src/autoinference_utils/cluster/modelexpress/dict_client.py +182 -0
  9. autoinference_utils-0.2.11/src/autoinference_utils/cluster/modelexpress/engine.py +130 -0
  10. autoinference_utils-0.2.11/src/autoinference_utils/cluster/modelexpress/image.py +102 -0
  11. autoinference_utils-0.2.11/src/autoinference_utils/cluster/modelexpress/peers.py +128 -0
  12. autoinference_utils-0.2.11/src/autoinference_utils/cluster/modelexpress/shim.py +81 -0
  13. {autoinference_utils-0.2.9 → autoinference_utils-0.2.11}/src/autoinference_utils/endpoint.py +87 -26
  14. {autoinference_utils-0.2.9 → autoinference_utils-0.2.11}/src/autoinference_utils/pd.py +75 -3
  15. {autoinference_utils-0.2.9 → autoinference_utils-0.2.11}/.gitignore +0 -0
  16. {autoinference_utils-0.2.9 → autoinference_utils-0.2.11}/README.md +0 -0
  17. {autoinference_utils-0.2.9 → autoinference_utils-0.2.11}/src/autoinference_utils/__init__.py +0 -0
  18. {autoinference_utils-0.2.9 → autoinference_utils-0.2.11}/src/autoinference_utils/router.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.5
2
2
  Name: autoinference-utils
3
- Version: 0.2.9
3
+ Version: 0.2.11
4
4
  Summary: Shared endpoint abstractions for autoinference deployments
5
5
  Requires-Python: >=3.10
6
6
  Description-Content-Type: text/markdown
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "autoinference-utils"
3
- version = "0.2.9"
3
+ version = "0.2.11"
4
4
  description = "Shared endpoint abstractions for autoinference deployments"
5
5
  readme = "README.md"
6
6
  requires-python = ">=3.10"
@@ -0,0 +1 @@
1
+ """Utilities for app-local inference clusters."""
@@ -0,0 +1,19 @@
1
+ from __future__ import annotations
2
+
3
+ import socket
4
+
5
+
6
+ def i6pn_address() -> str:
7
+ return socket.getaddrinfo("i6pn.modal.local", None, socket.AF_INET6)[0][4][0]
8
+
9
+
10
+ def endpoint(host: str, port: int) -> str:
11
+ return f"[{host}]:{port}" if ":" in host else f"{host}:{port}"
12
+
13
+
14
+ def dialable(address: str) -> str:
15
+ """Bracket bare IPv6 endpoints for gRPC."""
16
+ host, _, port = address.rpartition(":")
17
+ if not host or not port.isdigit():
18
+ return address
19
+ return endpoint(host.strip("[]"), int(port))
@@ -0,0 +1,11 @@
1
+ """Shared cluster discovery configuration."""
2
+
3
+ from __future__ import annotations
4
+
5
+ DICT_NAME_ENV = "AUTOINFERENCE_CLUSTER_DICT"
6
+ SCOPE_ENV = "AUTOINFERENCE_CLUSTER_SCOPE"
7
+ TTL_ENV = "AUTOINFERENCE_CLUSTER_TTL"
8
+
9
+ DEFAULT_DICT_NAME = "autoinf-cluster"
10
+
11
+ DEFAULT_TTL_SECONDS = 120.0
@@ -0,0 +1,126 @@
1
+ from __future__ import annotations
2
+
3
+ import modal
4
+
5
+ _ABSL_REPOSITORY = "https://github.com/abseil/abseil-cpp.git"
6
+ _UCX_REPOSITORY = "https://github.com/openucx/ucx.git"
7
+ _LIBFABRIC_REPOSITORY = "https://github.com/ofiwg/libfabric.git"
8
+
9
+
10
+ NIXL_REPOSITORY = "https://github.com/ai-dynamo/nixl.git"
11
+ NIXL_COMMIT = "26eabeb5667b386642298af948811ef77355d1d8"
12
+ # Upstream's version; Ubuntu's libabsl-dev lacks absl_log.
13
+ ABSL_TAG = "lts_2025_08_14"
14
+ LIBFABRIC_TAG = "v1.21.0"
15
+ # Match the UCX ABI already loaded by the engine images.
16
+ UCX_COMMIT = "5288e74cd40be622e109489de7bb80d821341bd5"
17
+
18
+ _NIXL_SOURCE_PATH = "/opt/nixl"
19
+
20
+ # UCX serves mlx5 hosts; libfabric serves EFA hosts.
21
+ _NIXL_MESON_ARGS = (
22
+ "-Csetup-args=-Denable_plugins=UCX,LIBFABRIC "
23
+ "-Csetup-args=-Ducx_path=/usr/local/ucx "
24
+ "-Csetup-args=-Dlibfabric_path=/usr/local "
25
+ "-Csetup-args=-Dbuild_tests=false "
26
+ "-Csetup-args=-Dbuild_examples=false "
27
+ "-Csetup-args=-Dbuild_docs=false"
28
+ )
29
+
30
+ # Catch stock libraries left behind under a different distribution name.
31
+ _VERIFY_NIXL_IPV6 = (
32
+ 'python -c "'
33
+ "import pathlib, nixl; "
34
+ "libs = sorted(pathlib.Path(nixl.__file__).parent.parent.rglob('libnixl*.so*')); "
35
+ "stock = [str(p) for p in libs if b'inet_pton failed for ip_addr' in p.read_bytes()]; "
36
+ "ipv6 = [str(p) for p in libs if b'Invalid IPv4 or IPv6 address: ' in p.read_bytes()]; "
37
+ "assert not stock, ('stock nixl still present: ' + str(stock)); "
38
+ "assert ipv6, ('NIXL lacks IPv6 support: ' + str([str(p) for p in libs])); "
39
+ "print('IPv6 NIXL:', ipv6)\""
40
+ )
41
+
42
+
43
+ def install_nixl(image: modal.Image) -> modal.Image:
44
+ """Build and install the pinned NIXL revision."""
45
+ return (
46
+ image.apt_install(
47
+ "git",
48
+ "cmake",
49
+ "pkg-config",
50
+ "ninja-build",
51
+ "autoconf",
52
+ "automake",
53
+ "libtool",
54
+ "libnuma-dev",
55
+ "libibverbs-dev",
56
+ "librdmacm-dev",
57
+ "ibverbs-providers",
58
+ )
59
+ .run_commands(
60
+ 'uv pip install --python "$(command -v python)" --no-cache '
61
+ "meson meson-python pybind11 patchelf build"
62
+ )
63
+ .run_commands(
64
+ "apt-get remove -y libabsl-dev || true",
65
+ "rm -rf /opt/absl && git init /opt/absl",
66
+ f"git -C /opt/absl remote add origin {_ABSL_REPOSITORY}",
67
+ f"git -C /opt/absl fetch --depth 1 origin {ABSL_TAG}",
68
+ "git -C /opt/absl checkout --detach FETCH_HEAD",
69
+ "cmake -S /opt/absl -B /opt/absl/build -GNinja "
70
+ "-DCMAKE_INSTALL_PREFIX=/usr/local -DCMAKE_BUILD_TYPE=Release "
71
+ "-DCMAKE_CXX_STANDARD=20 -DBUILD_SHARED_LIBS=ON "
72
+ "-DABSL_PROPAGATE_CXX_STD=ON -DABSL_ENABLE_INSTALL=ON",
73
+ "ninja -C /opt/absl/build && ninja -C /opt/absl/build install && ldconfig",
74
+ "rm -rf /opt/absl",
75
+ )
76
+ .run_commands(
77
+ "rm -rf /opt/ucx && git init /opt/ucx",
78
+ f"cd /opt/ucx && git remote add origin {_UCX_REPOSITORY} "
79
+ f"&& git fetch --depth 1 origin {UCX_COMMIT} "
80
+ "&& git checkout --detach FETCH_HEAD",
81
+ f'test "$(git -C /opt/ucx rev-parse HEAD)" = "{UCX_COMMIT}"',
82
+ # gVisor returns EINVAL for UCX's invalid-fd DMA-BUF probe.
83
+ "sed -i 's/errno == EBADF/errno == EBADF || errno == EINVAL/' "
84
+ "/opt/ucx/src/uct/ib/base/ib_md.c",
85
+ "grep -Fq 'errno == EBADF || errno == EINVAL' "
86
+ "/opt/ucx/src/uct/ib/base/ib_md.c",
87
+ "cd /opt/ucx && ./autogen.sh && ./contrib/configure-release-mt "
88
+ "--prefix=/usr/local/ucx --enable-shared --disable-static "
89
+ "--disable-doxygen-doc --enable-optimizations --enable-cma "
90
+ "--enable-devel-headers --with-cuda=/usr/local/cuda --with-verbs "
91
+ "--with-dm --without-gdrcopy",
92
+ "cd /opt/ucx && make -j$(nproc) && make install && ldconfig",
93
+ "rm -rf /opt/ucx",
94
+ )
95
+ .run_commands(
96
+ f"git clone --depth 1 -b {LIBFABRIC_TAG} {_LIBFABRIC_REPOSITORY} /opt/libfabric",
97
+ "cd /opt/libfabric && ./autogen.sh && ./configure --prefix=/usr/local "
98
+ "--disable-verbs --disable-psm3 --disable-opx --disable-usnic "
99
+ "--disable-rstream --enable-efa --with-cuda=/usr/local/cuda "
100
+ "--enable-cuda-dlopen --without-gdrcopy",
101
+ "cd /opt/libfabric && make -j$(nproc) && make install && ldconfig",
102
+ "rm -rf /opt/libfabric",
103
+ )
104
+ .run_commands(
105
+ f"rm -rf {_NIXL_SOURCE_PATH} && git init {_NIXL_SOURCE_PATH}",
106
+ f"cd {_NIXL_SOURCE_PATH} && git remote add origin {NIXL_REPOSITORY} "
107
+ f"&& git fetch --depth 1 origin {NIXL_COMMIT} "
108
+ "&& git checkout --detach FETCH_HEAD",
109
+ f'test "$(git -C {_NIXL_SOURCE_PATH} rev-parse HEAD)" = "{NIXL_COMMIT}"',
110
+ 'uv pip install --python "$(command -v python)" --no-cache tomlkit',
111
+ f"cd {_NIXL_SOURCE_PATH} && "
112
+ "nixl_cuda_major=$(nvcc --version | sed -n 's/.*release \\([0-9]*\\)\\..*/\\1/p') "
113
+ '&& python contrib/tomlutil.py --wheel-name "nixl-cu${nixl_cuda_major}" pyproject.toml',
114
+ 'uv pip uninstall --python "$(command -v python)" nixl nixl-cu12 nixl-cu13',
115
+ f'cd {_NIXL_SOURCE_PATH} && uv pip install --python "$(command -v python)" '
116
+ f"--no-cache --no-deps --no-build-isolation -Cbuild-dir={_NIXL_SOURCE_PATH}/build "
117
+ f"{_NIXL_MESON_ARGS} .",
118
+ 'uv pip install --python "$(command -v python)" --no-deps '
119
+ f"{_NIXL_SOURCE_PATH}/build/src/bindings/python/nixl-meta/nixl-*-py3-none-any.whl",
120
+ _VERIFY_NIXL_IPV6,
121
+ 'python -c "from nixl._api import nixl_agent, nixl_agent_config; '
122
+ "agent = nixl_agent('build-check', nixl_agent_config(backends=[])); "
123
+ "plugins = agent.get_plugin_list(); "
124
+ "assert {'UCX', 'LIBFABRIC'} <= set(plugins), plugins\"",
125
+ )
126
+ )
@@ -0,0 +1,18 @@
1
+ """ModelExpress weight-loading integration."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from typing import TYPE_CHECKING
6
+
7
+ if TYPE_CHECKING:
8
+ import modal
9
+
10
+ BACKEND_NAME = "modal-dict"
11
+
12
+
13
+ def install_cluster_deps(
14
+ image: "modal.Image", *, engine: str, enabled: bool = True
15
+ ) -> "modal.Image":
16
+ from .image import install_cluster_deps
17
+
18
+ return install_cluster_deps(image, engine=engine, enabled=enabled)
@@ -0,0 +1,182 @@
1
+ """ModelExpress worker discovery through Modal Dict."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from dataclasses import replace
6
+ import logging
7
+ import os
8
+ import threading
9
+ import time
10
+
11
+ from modelexpress import p2p_pb2
12
+ from modelexpress.client import MxClientBase
13
+ from modelexpress.metadata.source_id import compute_mx_source_id
14
+
15
+ from .._net import dialable
16
+ from ..config import (
17
+ DEFAULT_DICT_NAME,
18
+ DEFAULT_TTL_SECONDS,
19
+ DICT_NAME_ENV,
20
+ SCOPE_ENV,
21
+ TTL_ENV,
22
+ )
23
+ from .peers import PeerRecord, PeerStatus, PeerTable
24
+
25
+ logger = logging.getLogger("autoinference.cluster.modelexpress")
26
+
27
+
28
+ def _source_status(value: int) -> PeerStatus:
29
+ return PeerStatus[p2p_pb2.SourceStatus.Name(value).removeprefix("SOURCE_STATUS_")]
30
+
31
+
32
+ class MxModalDictClient(MxClientBase):
33
+ # Makes ModelExpress start WorkerGrpcServer for peer manifest retrieval.
34
+ REQUIRES_P2P_METADATA = True
35
+
36
+ def __init__(self, table: PeerTable):
37
+ self._table = table
38
+ self._published: dict[tuple[str, str], PeerRecord] = {}
39
+ self._lock = threading.Lock()
40
+ self._closed = False
41
+
42
+ def close(self) -> None:
43
+ with self._lock:
44
+ self._closed = True
45
+ for mx_source_id, worker_id in tuple(self._published):
46
+ try:
47
+ self._table.remove(mx_source_id, worker_id)
48
+ except Exception:
49
+ logger.warning(
50
+ "failed to remove source %s/%s",
51
+ mx_source_id,
52
+ worker_id,
53
+ exc_info=True,
54
+ )
55
+ else:
56
+ self._published.pop((mx_source_id, worker_id))
57
+
58
+ def publish_metadata(
59
+ self,
60
+ identity: "p2p_pb2.SourceIdentity",
61
+ worker: "p2p_pb2.WorkerMetadata",
62
+ worker_id: str,
63
+ ) -> str:
64
+ with self._lock:
65
+ if self._closed:
66
+ raise RuntimeError("cannot publish metadata after closing the client")
67
+ mx_source_id = compute_mx_source_id(identity)
68
+ record = PeerRecord(
69
+ group_id=mx_source_id,
70
+ worker_id=worker_id,
71
+ worker_rank=worker.worker_rank,
72
+ accelerator=worker.accelerator,
73
+ status=_source_status(
74
+ worker.status or p2p_pb2.SOURCE_STATUS_INITIALIZING
75
+ ),
76
+ grpc_endpoint=dialable(worker.worker_grpc_endpoint),
77
+ metadata_endpoint=worker.metadata_endpoint,
78
+ agent_name=worker.agent_name,
79
+ updated_at=time.time(),
80
+ )
81
+ self._table.put(record)
82
+ self._published[mx_source_id, worker_id] = record
83
+ logger.info(
84
+ "published source %s rank=%d at %s",
85
+ mx_source_id,
86
+ worker.worker_rank,
87
+ worker.worker_grpc_endpoint,
88
+ )
89
+ return mx_source_id
90
+
91
+ def list_sources(
92
+ self,
93
+ identity: "p2p_pb2.SourceIdentity | None" = None,
94
+ status_filter: "p2p_pb2.SourceStatus | None" = None,
95
+ ) -> "p2p_pb2.ListSourcesResponse":
96
+ if identity is None:
97
+ raise ValueError(
98
+ "list_sources requires an identity so mx_source_id can be "
99
+ "computed locally without a central coordinator"
100
+ )
101
+ mx_source_id = compute_mx_source_id(identity)
102
+ expected_status = (
103
+ _source_status(status_filter) if status_filter is not None else None
104
+ )
105
+ instances = [
106
+ p2p_pb2.SourceInstanceRef(
107
+ mx_source_id=record.group_id,
108
+ worker_id=record.worker_id,
109
+ model_name=identity.model_name,
110
+ worker_rank=record.worker_rank,
111
+ accelerator=record.accelerator,
112
+ )
113
+ for record in self._table.list(mx_source_id)
114
+ if record.grpc_endpoint
115
+ and (expected_status is None or record.status == expected_status)
116
+ ]
117
+ logger.info("listed %d source(s) for %s", len(instances), mx_source_id)
118
+ return p2p_pb2.ListSourcesResponse(instances=instances)
119
+
120
+ def get_metadata(
121
+ self,
122
+ mx_source_id: str,
123
+ worker_id: str,
124
+ ) -> "p2p_pb2.GetMetadataResponse":
125
+ record = self._table.get(mx_source_id, worker_id)
126
+ if record is None or not record.grpc_endpoint:
127
+ logger.warning("no live record for %s/%s", mx_source_id, worker_id)
128
+ return p2p_pb2.GetMetadataResponse(
129
+ found=False, mx_source_id=mx_source_id, worker_id=worker_id
130
+ )
131
+
132
+ return p2p_pb2.GetMetadataResponse(
133
+ found=True,
134
+ mx_source_id=mx_source_id,
135
+ worker_id=worker_id,
136
+ worker=p2p_pb2.WorkerMetadata(
137
+ worker_rank=record.worker_rank,
138
+ accelerator=record.accelerator,
139
+ status=p2p_pb2.SourceStatus.Value(
140
+ f"SOURCE_STATUS_{record.status.name}"
141
+ ),
142
+ worker_grpc_endpoint=record.grpc_endpoint,
143
+ metadata_endpoint=record.metadata_endpoint,
144
+ agent_name=record.agent_name,
145
+ ),
146
+ )
147
+
148
+ def update_status(
149
+ self,
150
+ mx_source_id: str,
151
+ worker_id: str,
152
+ worker_rank: int,
153
+ status: "p2p_pb2.SourceStatus",
154
+ source_load: float | None = None,
155
+ ) -> bool:
156
+ with self._lock:
157
+ record = self._published.get((mx_source_id, worker_id))
158
+ if record is None or self._closed:
159
+ logger.warning(
160
+ "cannot update status for unknown %s/%s", mx_source_id, worker_id
161
+ )
162
+ return False
163
+ record = replace(
164
+ record, status=_source_status(status), updated_at=time.time()
165
+ )
166
+ self._table.put(record)
167
+ self._published[mx_source_id, worker_id] = record
168
+ return True
169
+
170
+
171
+ def build_dict_client() -> MxModalDictClient:
172
+ scope = os.environ.get(SCOPE_ENV, "").strip()
173
+ if not scope:
174
+ raise RuntimeError(
175
+ f"{SCOPE_ENV} is unset; the cluster scope must reach the engine"
176
+ )
177
+ table = PeerTable(
178
+ scope=scope,
179
+ dict_name=os.environ.get(DICT_NAME_ENV) or DEFAULT_DICT_NAME,
180
+ ttl=float(os.environ.get(TTL_ENV) or DEFAULT_TTL_SECONDS),
181
+ )
182
+ return MxModalDictClient(table)
@@ -0,0 +1,130 @@
1
+ """ModelExpress engine configuration for peer weight loading."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import os
6
+ from dataclasses import dataclass, field
7
+ from pathlib import Path
8
+ from typing import Mapping, Optional
9
+
10
+ from .._net import i6pn_address
11
+ from ..config import (
12
+ DEFAULT_DICT_NAME,
13
+ DEFAULT_TTL_SECONDS,
14
+ DICT_NAME_ENV,
15
+ SCOPE_ENV,
16
+ TTL_ENV,
17
+ )
18
+
19
+ from . import BACKEND_NAME
20
+
21
+ # Bound upstream's transfer and publish waits for cold deployments.
22
+ DEFAULT_TRANSFER_TIMEOUT = 120
23
+ DEFAULT_PUBLISH_TIMEOUT = 120
24
+
25
+ # UCX auto-detection includes the CUDA transports needed to register weights.
26
+ DEFAULT_NIXL_UCX_TLS = None
27
+
28
+ # SGLang and ModelExpress require a non-empty URL; the Dict client ignores it.
29
+ MX_PLACEHOLDER_URL = "modal-dict://peers"
30
+ MX_CONFIG = '{"transport":"nixl","url":"%s"}' % MX_PLACEHOLDER_URL
31
+
32
+ _SGLANG_ARGS = {
33
+ "--load-format": "remote_instance",
34
+ "--remote-instance-weight-loader-backend": "modelexpress",
35
+ "--modelexpress-config": MX_CONFIG,
36
+ # RemoteInstanceModelLoader rejects SGLangEndpoint's default extra config.
37
+ "--model-loader-extra-config": "{}",
38
+ }
39
+ _VLLM_ARGS = {"--load-format": "modelexpress"}
40
+
41
+
42
+ @dataclass(frozen=True)
43
+ class _EngineOverrides:
44
+ """Launch overrides for one engine process. Falsy when the cluster is off."""
45
+
46
+ scope: Optional[str] = None
47
+ server_args: Mapping[str, str] = field(default_factory=dict)
48
+ env: Mapping[str, str] = field(default_factory=dict)
49
+
50
+ def __bool__(self) -> bool:
51
+ return self.scope is not None
52
+
53
+
54
+ def _engine_overrides(
55
+ *,
56
+ app_name: str,
57
+ app_id: str,
58
+ engine: str,
59
+ revision: str = "",
60
+ dict_name: Optional[str] = None,
61
+ ttl: float = DEFAULT_TTL_SECONDS,
62
+ nixl_ucx_tls: Optional[str] = DEFAULT_NIXL_UCX_TLS,
63
+ transfer_timeout: int = DEFAULT_TRANSFER_TIMEOUT,
64
+ publish_timeout: int = DEFAULT_PUBLISH_TIMEOUT,
65
+ ) -> _EngineOverrides:
66
+ """Render the overrides for an engine launched from this container."""
67
+ if engine not in ("sglang", "vllm"):
68
+ raise ValueError(f"unsupported engine {engine!r}")
69
+ if not app_id:
70
+ raise ValueError(
71
+ "app_id is unset; configure the cluster from a running Modal app"
72
+ )
73
+
74
+ region = os.environ.get("MODAL_REGION", "")
75
+ if not region:
76
+ return _off("MODAL_REGION is unset")
77
+ try:
78
+ worker_host = i6pn_address()
79
+ except OSError as error:
80
+ return _off(f"no i6pn address: {error}")
81
+
82
+ scope = f"weights/{app_id}/{region}"
83
+ # Modal limits Dict names to 64 characters.
84
+ dict_prefix = f"{DEFAULT_DICT_NAME}-{app_name}"
85
+ dict_name = dict_name or f"{dict_prefix[:55]}-{app_id[-8:]}"
86
+ env = {
87
+ "MX_METADATA_BACKEND": BACKEND_NAME,
88
+ "MX_SERVER_ADDRESS": MX_PLACEHOLDER_URL,
89
+ SCOPE_ENV: scope,
90
+ DICT_NAME_ENV: dict_name,
91
+ TTL_ENV: str(ttl),
92
+ # NIXL needs a bare host; the ModelExpress adapter brackets gRPC endpoints.
93
+ "MX_WORKER_HOST": worker_host,
94
+ "MX_NIXL_BACKEND": _nixl_backend(),
95
+ "MX_P2P_METADATA": "1",
96
+ "MX_P2P_SOURCE_SELECTOR": "random",
97
+ # i6pn is IPv6; container IPv4 addresses are not peer-routable.
98
+ "UCX_TCP_AF_PRIO": "inet6",
99
+ "MX_NIXL_METADATA_TIMEOUT": "5",
100
+ "MX_TRANSFER_TIMEOUT": str(transfer_timeout),
101
+ "MX_PUBLISH_TIMEOUT_SECS": str(publish_timeout),
102
+ # SGLang's multimodal workers fork after gRPC has started.
103
+ "GRPC_ENABLE_FORK_SUPPORT": "1",
104
+ "GRPC_POLL_STRATEGY": "poll",
105
+ }
106
+ if revision:
107
+ env["MX_MODEL_REVISION"] = revision
108
+ if nixl_ucx_tls:
109
+ env["NIXL_UCX_TLS"] = nixl_ucx_tls
110
+
111
+ print(f"[cluster] enabled, scope={scope} host={worker_host}", flush=True)
112
+ return _EngineOverrides(
113
+ scope=scope,
114
+ server_args=dict(_SGLANG_ARGS if engine == "sglang" else _VLLM_ARGS),
115
+ env=env,
116
+ )
117
+
118
+
119
+ def _nixl_backend() -> str:
120
+ """Use libfabric for mounted EFA devices; UCX handles verbs and TCP."""
121
+ for device in Path("/dev/infiniband").glob("uverbs*"):
122
+ driver = Path("/sys/class/infiniband_verbs") / device.name / "device" / "driver"
123
+ if driver.is_symlink() and os.path.basename(os.readlink(driver)) == "efa":
124
+ return "LIBFABRIC"
125
+ return "UCX"
126
+
127
+
128
+ def _off(reason: str) -> _EngineOverrides:
129
+ print(f"[cluster] disabled, loading weights normally: {reason}", flush=True)
130
+ return _EngineOverrides()
@@ -0,0 +1,102 @@
1
+ from __future__ import annotations
2
+
3
+ import base64
4
+
5
+ import modal
6
+
7
+ from ..image import install_nixl
8
+
9
+ MODELEXPRESS_REPOSITORY = "https://github.com/modal-projects/modelexpress.git"
10
+ MODELEXPRESS_COMMIT = "7fc747444459423029a37bb279b9bc6d401d9d42"
11
+ MODELEXPRESS_VERSION = "0.7.0"
12
+ _SOURCE_PATH = "/opt/modelexpress"
13
+
14
+ # Run after other .pth files have adjusted sys.path.
15
+ _PTH_NAME = "zz-autoinference-cluster.pth"
16
+
17
+ # .pth files execute single lines starting with "import ". Guard optional imports.
18
+ _PTH_LINE = (
19
+ r"import sys; exec('try:\n"
20
+ r" import autoinference_utils.cluster.modelexpress.shim\n"
21
+ r"except Exception:\n"
22
+ r" pass\n')"
23
+ )
24
+ _PTH_B64 = base64.b64encode(f"{_PTH_LINE}\n".encode()).decode()
25
+
26
+ # Co-locate with the package: pip and uv can select different site directories.
27
+ _WRITE_PTH = (
28
+ 'python -c "import base64, pathlib, autoinference_utils; '
29
+ "pathlib.Path(autoinference_utils.__file__).parent.parent"
30
+ f".joinpath('{_PTH_NAME}').write_bytes(base64.b64decode('{_PTH_B64}'))\""
31
+ )
32
+
33
+ # The startup hook must work without Modal's injected PYTHONPATH.
34
+ _VERIFY_SHIM = (
35
+ 'env -u PYTHONPATH MX_METADATA_BACKEND=modal-dict python -c "'
36
+ "import sys; "
37
+ "names = [type(f).__name__ for f in sys.meta_path]; "
38
+ "assert '_PatchOnImport' in names, names\""
39
+ )
40
+
41
+ _VERIFY_SHIM_IS_SILENT = (
42
+ "test -z \"$(env -u PYTHONPATH -u MX_METADATA_BACKEND python -c 'pass' 2>&1)\""
43
+ )
44
+
45
+ # SGLang #24723 replaced inline ModelExpress loading with package delegation.
46
+ SGLANG_MIN_COMMIT = "435ea41cf0"
47
+ SGLANG_MIN_DATE = "2026-05-16"
48
+
49
+ _ENGINE_DELEGATION_CHECKS = {
50
+ "sglang": (
51
+ 'python -c "'
52
+ "import inspect; from sglang.srt.model_loader import loader; "
53
+ "assert 'modelexpress.engines.sglang.loader' in inspect.getsource(loader), "
54
+ "'This SGLang build predates #24723 (%s, %s): it builds its own "
55
+ "MxClient and ignores the modal-dict backend. Use a newer image.'\""
56
+ % (SGLANG_MIN_COMMIT, SGLANG_MIN_DATE)
57
+ ),
58
+ "vllm": (
59
+ 'python -c "'
60
+ "from vllm.model_executor.model_loader import _LOAD_FORMAT_TO_MODEL_LOADER as m; "
61
+ "assert 'modelexpress' in m, sorted(m)\""
62
+ ),
63
+ }
64
+
65
+ _VERIFY_MODELEXPRESS = (
66
+ 'python -c "from importlib.metadata import version; '
67
+ f"assert version('modelexpress') == '{MODELEXPRESS_VERSION}'; "
68
+ "import modelexpress.nixl_transfer; "
69
+ 'from modelexpress.metadata.client_factory import create_metadata_client"'
70
+ )
71
+
72
+
73
+ def install_cluster_deps(
74
+ image: modal.Image,
75
+ *,
76
+ engine: str,
77
+ enabled: bool = True,
78
+ ) -> modal.Image:
79
+ """Add NIXL and ModelExpress peer loading to an engine image.
80
+
81
+ autoinference-utils and runtime dependencies must already be installed.
82
+ Use --no-deps to preserve the engine image's torch installation.
83
+ """
84
+ if not enabled:
85
+ return image
86
+ if engine not in _ENGINE_DELEGATION_CHECKS:
87
+ raise ValueError(f"unsupported engine {engine!r}")
88
+ return install_nixl(image).run_commands(
89
+ f"rm -rf {_SOURCE_PATH} && git init {_SOURCE_PATH}",
90
+ f"git -C {_SOURCE_PATH} remote add origin {MODELEXPRESS_REPOSITORY}",
91
+ f"git -C {_SOURCE_PATH} fetch --depth 1 origin {MODELEXPRESS_COMMIT}",
92
+ f"git -C {_SOURCE_PATH} checkout --detach FETCH_HEAD",
93
+ f'test "$(git -C {_SOURCE_PATH} rev-parse HEAD)" = "{MODELEXPRESS_COMMIT}"',
94
+ # --system can select a different interpreter in engine images.
95
+ 'uv pip install --python "$(command -v python)" --no-cache --no-deps '
96
+ f"-e {_SOURCE_PATH}/modelexpress_client/python",
97
+ _VERIFY_MODELEXPRESS,
98
+ _WRITE_PTH,
99
+ _VERIFY_SHIM,
100
+ _VERIFY_SHIM_IS_SILENT,
101
+ _ENGINE_DELEGATION_CHECKS[engine],
102
+ )
@@ -0,0 +1,128 @@
1
+ from __future__ import annotations
2
+
3
+ import time
4
+ from dataclasses import dataclass
5
+ from enum import Enum
6
+ from typing import Any, Optional
7
+
8
+ import modal
9
+
10
+ from ..config import DEFAULT_DICT_NAME, DEFAULT_TTL_SECONDS
11
+
12
+
13
+ class PeerStatus(str, Enum):
14
+ UNKNOWN = "unknown"
15
+ INITIALIZING = "initializing"
16
+ READY = "ready"
17
+ STALE = "stale"
18
+
19
+
20
+ @dataclass(frozen=True)
21
+ class PeerRecord:
22
+ group_id: str
23
+ worker_id: str
24
+ worker_rank: int
25
+ accelerator: str
26
+ status: PeerStatus
27
+ grpc_endpoint: str
28
+ metadata_endpoint: str
29
+ agent_name: str
30
+ updated_at: float
31
+
32
+ def as_payload(self) -> dict[str, Any]:
33
+ return {
34
+ "worker_rank": self.worker_rank,
35
+ "accelerator": self.accelerator,
36
+ "status": self.status.value,
37
+ "grpc_endpoint": self.grpc_endpoint,
38
+ "metadata_endpoint": self.metadata_endpoint,
39
+ "agent_name": self.agent_name,
40
+ "updated_at": self.updated_at,
41
+ }
42
+
43
+ @classmethod
44
+ def from_payload(
45
+ cls, raw: Any, *, group_id: str, worker_id: str
46
+ ) -> Optional["PeerRecord"]:
47
+ if not isinstance(raw, dict):
48
+ return None
49
+ try:
50
+ return cls(
51
+ group_id=group_id,
52
+ worker_id=worker_id,
53
+ worker_rank=int(raw["worker_rank"]),
54
+ accelerator=str(raw["accelerator"]),
55
+ status=PeerStatus(raw["status"]),
56
+ grpc_endpoint=str(raw["grpc_endpoint"]),
57
+ metadata_endpoint=str(raw["metadata_endpoint"]),
58
+ agent_name=str(raw["agent_name"]),
59
+ updated_at=float(raw["updated_at"]),
60
+ )
61
+ except (KeyError, TypeError, ValueError):
62
+ return None
63
+
64
+
65
+ class PeerTable:
66
+ def __init__(
67
+ self,
68
+ *,
69
+ scope: str,
70
+ dict_name: str = DEFAULT_DICT_NAME,
71
+ ttl: float = DEFAULT_TTL_SECONDS,
72
+ registry: Optional[modal.Dict] = None,
73
+ ):
74
+ if not scope:
75
+ raise ValueError("a peer table needs a scope")
76
+ self.scope = scope
77
+ self.ttl = ttl
78
+ self._registry = registry or modal.Dict.from_name(
79
+ dict_name, create_if_missing=True
80
+ )
81
+
82
+ def key(self, group_id: str, worker_id: str) -> str:
83
+ if not worker_id or "/" in worker_id:
84
+ raise ValueError("worker_id must be non-empty and cannot contain '/'")
85
+ return f"{self.scope}/{group_id}/{worker_id}"
86
+
87
+ def put(self, record: PeerRecord) -> None:
88
+ self._registry.put(
89
+ self.key(record.group_id, record.worker_id),
90
+ record.as_payload(),
91
+ )
92
+
93
+ def get(self, group_id: str, worker_id: str) -> Optional[PeerRecord]:
94
+ raw = self._registry.get(self.key(group_id, worker_id))
95
+ if raw is None:
96
+ return None
97
+ record = PeerRecord.from_payload(raw, group_id=group_id, worker_id=worker_id)
98
+ cutoff = time.time() - self.ttl
99
+ if record is None or record.updated_at < cutoff:
100
+ self._reap(self.key(group_id, worker_id), group_id, worker_id, cutoff)
101
+ return None
102
+ return record
103
+
104
+ def list(self, group_id: str) -> list[PeerRecord]:
105
+ prefix = f"{self.scope}/"
106
+ cutoff = time.time() - self.ttl
107
+ live: list[PeerRecord] = []
108
+ for key, raw in self._registry.items():
109
+ if not isinstance(key, str) or not key.startswith(prefix):
110
+ continue
111
+ group, _, worker_id = key[len(prefix) :].partition("/")
112
+ if not group or not worker_id or "/" in worker_id:
113
+ continue
114
+ record = PeerRecord.from_payload(raw, group_id=group, worker_id=worker_id)
115
+ if record is None or record.updated_at < cutoff:
116
+ self._reap(key, group, worker_id, cutoff)
117
+ elif group == group_id:
118
+ live.append(record)
119
+ return live
120
+
121
+ def _reap(self, key: str, group_id: str, worker_id: str, cutoff: float) -> None:
122
+ raw = self._registry.pop(key, None)
123
+ record = PeerRecord.from_payload(raw, group_id=group_id, worker_id=worker_id)
124
+ if record is not None and record.updated_at >= cutoff:
125
+ self._registry.put(key, raw, skip_if_exists=True)
126
+
127
+ def remove(self, group_id: str, worker_id: str) -> None:
128
+ self._registry.pop(self.key(group_id, worker_id), None)
@@ -0,0 +1,81 @@
1
+ """Register the Dict backend through a .pth import hook.
2
+
3
+ ModelExpress has no backend registration API. Defer its import until the engine
4
+ loads it so the startup hook does not import torch.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ import importlib.util
10
+ import os
11
+ import sys
12
+ from importlib.abc import Loader, MetaPathFinder
13
+ from typing import Any, Optional
14
+
15
+ from . import BACKEND_NAME
16
+
17
+ TARGET_MODULE = "modelexpress.metadata.client_factory"
18
+
19
+
20
+ def install() -> None:
21
+ if os.environ.get("MX_METADATA_BACKEND", "").strip().lower() != BACKEND_NAME:
22
+ return
23
+ if any(isinstance(finder, _PatchOnImport) for finder in sys.meta_path):
24
+ return
25
+ sys.meta_path.insert(0, _PatchOnImport())
26
+
27
+
28
+ class _PatchOnImport(MetaPathFinder):
29
+ def find_spec(self, fullname: str, path: Any = None, target: Any = None) -> Any:
30
+ if fullname != TARGET_MODULE:
31
+ return None
32
+ sys.meta_path.remove(self)
33
+ try:
34
+ spec = importlib.util.find_spec(fullname)
35
+ except Exception:
36
+ spec = None
37
+ finally:
38
+ sys.meta_path.insert(0, self)
39
+ if spec is None or spec.loader is None:
40
+ return None
41
+ spec.loader = _PatchingLoader(spec.loader)
42
+ return spec
43
+
44
+
45
+ class _PatchingLoader(Loader):
46
+ def __init__(self, inner: Loader):
47
+ self._inner = inner
48
+
49
+ def create_module(self, spec: Any) -> Any:
50
+ return self._inner.create_module(spec) # type: ignore[attr-defined]
51
+
52
+ def exec_module(self, module: Any) -> None:
53
+ self._inner.exec_module(module) # type: ignore[attr-defined]
54
+ patch_factory(module)
55
+
56
+
57
+ def patch_factory(module: Any) -> None:
58
+ original = getattr(module, "create_metadata_client", None)
59
+ if original is None or getattr(original, "_cluster_patched", False):
60
+ return
61
+
62
+ def create_metadata_client(
63
+ worker_rank: Optional[int] = None,
64
+ server_url: Optional[str] = None,
65
+ ) -> Any:
66
+ backend = os.environ.get("MX_METADATA_BACKEND", "").strip().lower()
67
+ if backend != BACKEND_NAME:
68
+ return original(worker_rank, server_url)
69
+ from .dict_client import build_dict_client
70
+
71
+ return build_dict_client()
72
+
73
+ create_metadata_client._cluster_patched = True # type: ignore[attr-defined]
74
+ module.create_metadata_client = create_metadata_client
75
+
76
+
77
+ try:
78
+ install()
79
+ except Exception: # noqa: BLE001
80
+ # .pth errors would be printed by every interpreter using this image.
81
+ pass
@@ -90,6 +90,59 @@ class Endpoint(ABC):
90
90
 
91
91
  def __init__(self, base_url: str):
92
92
  self.base_url = base_url.rstrip("/")
93
+ self._cluster_env: dict[str, str] = {}
94
+
95
+ def with_cluster(self, app, *, revision: str = "", enabled: bool = True):
96
+ """Configure peer weight loading before starting the endpoint."""
97
+ if not enabled:
98
+ return self
99
+ from .cluster.modelexpress.engine import _engine_overrides
100
+
101
+ if isinstance(self, SGLangEndpoint):
102
+ engine = "sglang"
103
+ elif isinstance(self, VLLMEndpoint):
104
+ engine = "vllm"
105
+ else:
106
+ raise ValueError(f"cluster is unsupported for {type(self).__name__}")
107
+ if self._proc is not None:
108
+ raise RuntimeError("configure the cluster before starting the endpoint")
109
+ overrides = _engine_overrides(
110
+ app_name=app.name, app_id=app.app_id, engine=engine, revision=revision
111
+ )
112
+ if not overrides:
113
+ return self
114
+ loader_args = {}
115
+ if isinstance(self, SGLangEndpoint) and self.load_format is not None:
116
+ loader_args["--load-format"] = self.load_format
117
+ loader_args = _merge_server_args(loader_args, self.extra_server_args)
118
+ for key, value in loader_args.items():
119
+ flag, attached, inline_value = key.partition("=")
120
+ value = inline_value if attached else value
121
+ if flag == "--load-format":
122
+ supported = value in ("auto", overrides.server_args[flag])
123
+ elif flag == "--model-loader-extra-config":
124
+ supported_configs = [{}]
125
+ if isinstance(self, SGLangEndpoint):
126
+ supported_configs.append(
127
+ json.loads(self.DEFAULT_OPERATIONAL_ARGS[flag])
128
+ )
129
+ try:
130
+ supported = json.loads(value) in supported_configs
131
+ except ValueError:
132
+ supported = False
133
+ elif flag in overrides.server_args:
134
+ supported = value == overrides.server_args[flag]
135
+ else:
136
+ continue
137
+ if not supported:
138
+ raise ValueError(f"cluster does not support {flag}={value!r}")
139
+ self.extra_server_args = _merge_server_args(
140
+ self.extra_server_args, overrides.server_args
141
+ )
142
+ self._cluster_env = dict(overrides.env)
143
+ if isinstance(self, SGLangEndpoint):
144
+ self.load_format = None
145
+ return self
93
146
 
94
147
  def __enter__(self):
95
148
  try:
@@ -300,21 +353,25 @@ class SGLangEndpoint(Endpoint):
300
353
  def start(self):
301
354
  cmd = self._build_cmd()
302
355
  print(f"[endpoint] starting: {shlex.join(cmd)}")
303
- self._proc = subprocess.Popen(cmd)
304
- wait_ready(
305
- self._proc,
306
- port=self.worker_port,
307
- base_url=self.base_url,
308
- timeout=self.health_timeout,
309
- poll_interval=self.health_poll_interval,
310
- request_timeout=self.health_request_timeout,
311
- )
312
- if self.bench_mode:
313
- self._bench_server = start_bench_proxy(
314
- listen_port=self.listen_port,
315
- upstream_port=self.worker_port,
316
- upstream_base_url=self.base_url,
356
+ self._proc = subprocess.Popen(cmd, env=os.environ | self._cluster_env)
357
+ try:
358
+ wait_ready(
359
+ self._proc,
360
+ port=self.worker_port,
361
+ base_url=self.base_url,
362
+ timeout=self.health_timeout,
363
+ poll_interval=self.health_poll_interval,
364
+ request_timeout=self.health_request_timeout,
317
365
  )
366
+ if self.bench_mode:
367
+ self._bench_server = start_bench_proxy(
368
+ listen_port=self.listen_port,
369
+ upstream_port=self.worker_port,
370
+ upstream_base_url=self.base_url,
371
+ )
372
+ except BaseException:
373
+ self.stop()
374
+ raise
318
375
 
319
376
  def stop(self):
320
377
  if self._bench_server is not None:
@@ -378,19 +435,23 @@ class VLLMEndpoint(Endpoint):
378
435
  def start(self):
379
436
  cmd = self._build_cmd()
380
437
  print(f"[vllm] starting: {shlex.join(cmd)}")
381
- self._proc = subprocess.Popen(cmd)
382
- wait_ready(
383
- self._proc,
384
- port=self.worker_port,
385
- timeout=self.health_timeout,
386
- poll_interval=self.health_poll_interval,
387
- request_timeout=self.health_request_timeout,
388
- )
389
- if self.bench_mode:
390
- self._bench_server = start_bench_proxy(
391
- listen_port=self.listen_port,
392
- upstream_port=self.worker_port,
438
+ self._proc = subprocess.Popen(cmd, env=os.environ | self._cluster_env)
439
+ try:
440
+ wait_ready(
441
+ self._proc,
442
+ port=self.worker_port,
443
+ timeout=self.health_timeout,
444
+ poll_interval=self.health_poll_interval,
445
+ request_timeout=self.health_request_timeout,
393
446
  )
447
+ if self.bench_mode:
448
+ self._bench_server = start_bench_proxy(
449
+ listen_port=self.listen_port,
450
+ upstream_port=self.worker_port,
451
+ )
452
+ except BaseException:
453
+ self.stop()
454
+ raise
394
455
 
395
456
  def stop(self):
396
457
  if self._bench_server is not None:
@@ -5,6 +5,8 @@ from __future__ import annotations
5
5
  import asyncio
6
6
  import math
7
7
  import os
8
+ import socket
9
+ import struct
8
10
  import subprocess
9
11
  import threading
10
12
  import time
@@ -40,6 +42,9 @@ _MANAGED_ROUTER_FLAGS = frozenset(
40
42
  }
41
43
  )
42
44
 
45
+ # Leave headroom inside Modal's 30-second exit-handler deadline.
46
+ _SHUTDOWN_TIMEOUT = 25
47
+
43
48
 
44
49
  def _managed_router_tokens(router_args: Mapping[str, str]) -> list[str]:
45
50
  """Rendered override tokens that would reach a managed flag.
@@ -177,6 +182,7 @@ class PDEndpoint(Endpoint):
177
182
  health_failure_timeout: float = 30,
178
183
  router_args: Mapping[str, str] | None = None,
179
184
  allow_custom_router: bool = False,
185
+ shutdown_port: int | None = None,
180
186
  ):
181
187
  from modal.experimental import get_cluster_info
182
188
 
@@ -217,6 +223,13 @@ class PDEndpoint(Endpoint):
217
223
  f"P/D ratio {ratio} requires a {len(roles)}-container Modal cluster"
218
224
  )
219
225
  self._host_ip = hosts[cluster.rank]
226
+ self._hosts = hosts
227
+ self._rank = cluster.rank
228
+ self._shutdown_port = (
229
+ router_port + 1 if shutdown_port is None else shutdown_port
230
+ )
231
+ if not 1 <= self._shutdown_port <= 65535:
232
+ raise ValueError("shutdown_port must be between 1 and 65535")
220
233
  self.role = roles[cluster.rank]
221
234
  super().__init__(_url(hosts[0], router_port))
222
235
  self.engine = engine
@@ -266,6 +279,8 @@ class PDEndpoint(Endpoint):
266
279
  self._stopped = threading.Event()
267
280
  self._lock = threading.Lock()
268
281
  self._started = False
282
+ self._serving = False
283
+ self._shutdown_listener: socket.socket | None = None
269
284
  self._threads: list[threading.Thread] = []
270
285
 
271
286
  def start(self, *, warmup: Callable[[], None] | None = None) -> None:
@@ -280,6 +295,13 @@ class PDEndpoint(Endpoint):
280
295
  SGLANG_ENABLE_HEALTH_ENDPOINT_GENERATION="1",
281
296
  )
282
297
  try:
298
+ if self._rank == 0:
299
+ self._shutdown_listener = socket.socket(socket.AF_INET6)
300
+ self._shutdown_listener.setsockopt(
301
+ socket.SOL_SOCKET, socket.SO_REUSEADDR, 1
302
+ )
303
+ self._shutdown_listener.bind(("::", self._shutdown_port))
304
+ self._shutdown_listener.listen(len(self._hosts) - 1)
283
305
  for endpoint in (self.engine, self.router):
284
306
  if endpoint is not None:
285
307
  endpoint.start()
@@ -291,6 +313,7 @@ class PDEndpoint(Endpoint):
291
313
  for _, host in self.router.pd_config:
292
314
  self._watch_health(_url(host, self.engine.worker_port) + "/health")
293
315
  self._watch_health(self.router.base_url + "/health")
316
+ self._serving = True
294
317
  except BaseException:
295
318
  self.stop()
296
319
  raise
@@ -352,7 +375,7 @@ class PDEndpoint(Endpoint):
352
375
  if self._stopped.is_set():
353
376
  return
354
377
  self._stopped.set()
355
- deadline = time.monotonic() + self.drain_timeout + 30
378
+ deadline = time.monotonic() + _SHUTDOWN_TIMEOUT
356
379
  try:
357
380
  for endpoint in (self.router, self.engine):
358
381
  if endpoint is None:
@@ -364,9 +387,58 @@ class PDEndpoint(Endpoint):
364
387
  process.wait(timeout=max(0, deadline - time.monotonic()))
365
388
  except subprocess.TimeoutExpired:
366
389
  process.kill()
390
+ process.wait()
367
391
  endpoint.stop()
368
392
  finally:
369
393
  for thread in self._threads:
370
- thread.join(
371
- timeout=self.engine.health_request_timeout + self.health_interval
394
+ thread.join(timeout=max(0, deadline - time.monotonic()))
395
+ if self._serving:
396
+ if self._rank == 0:
397
+ self._wait_for_followers(deadline)
398
+ else:
399
+ self._notify_leader(deadline)
400
+ if self._shutdown_listener is not None:
401
+ self._shutdown_listener.close()
402
+ self._shutdown_listener = None
403
+
404
+ def _notify_leader(self, deadline: float) -> None:
405
+ while True:
406
+ remaining = deadline - time.monotonic()
407
+ if remaining <= 0:
408
+ print("[pd] Failed to notify rank zero of shutdown", flush=True)
409
+ return
410
+ try:
411
+ with socket.create_connection(
412
+ (self._hosts[0], self._shutdown_port), timeout=min(remaining, 1)
413
+ ) as connection:
414
+ connection.sendall(struct.pack("!I", self._rank))
415
+ return
416
+ except OSError:
417
+ time.sleep(min(0.1, max(0, deadline - time.monotonic())))
418
+
419
+ def _wait_for_followers(self, deadline: float) -> None:
420
+ listener = self._shutdown_listener
421
+ if listener is None:
422
+ return
423
+ pending = set(range(1, len(self._hosts)))
424
+ while pending:
425
+ remaining = deadline - time.monotonic()
426
+ if remaining <= 0:
427
+ print(
428
+ f"[pd] Timed out waiting for ranks {sorted(pending)} to stop",
429
+ flush=True,
372
430
  )
431
+ return
432
+ listener.settimeout(min(remaining, 1))
433
+ try:
434
+ connection, _ = listener.accept()
435
+ except TimeoutError:
436
+ continue
437
+ with connection:
438
+ connection.settimeout(min(remaining, 1))
439
+ try:
440
+ payload = connection.recv(4, socket.MSG_WAITALL)
441
+ except TimeoutError:
442
+ continue
443
+ if len(payload) == 4:
444
+ pending.discard(struct.unpack("!I", payload)[0])