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.
- {autoinference_utils-0.2.9 → autoinference_utils-0.2.11}/PKG-INFO +1 -1
- {autoinference_utils-0.2.9 → autoinference_utils-0.2.11}/pyproject.toml +1 -1
- autoinference_utils-0.2.11/src/autoinference_utils/cluster/__init__.py +1 -0
- autoinference_utils-0.2.11/src/autoinference_utils/cluster/_net.py +19 -0
- autoinference_utils-0.2.11/src/autoinference_utils/cluster/config.py +11 -0
- autoinference_utils-0.2.11/src/autoinference_utils/cluster/image.py +126 -0
- autoinference_utils-0.2.11/src/autoinference_utils/cluster/modelexpress/__init__.py +18 -0
- autoinference_utils-0.2.11/src/autoinference_utils/cluster/modelexpress/dict_client.py +182 -0
- autoinference_utils-0.2.11/src/autoinference_utils/cluster/modelexpress/engine.py +130 -0
- autoinference_utils-0.2.11/src/autoinference_utils/cluster/modelexpress/image.py +102 -0
- autoinference_utils-0.2.11/src/autoinference_utils/cluster/modelexpress/peers.py +128 -0
- autoinference_utils-0.2.11/src/autoinference_utils/cluster/modelexpress/shim.py +81 -0
- {autoinference_utils-0.2.9 → autoinference_utils-0.2.11}/src/autoinference_utils/endpoint.py +87 -26
- {autoinference_utils-0.2.9 → autoinference_utils-0.2.11}/src/autoinference_utils/pd.py +75 -3
- {autoinference_utils-0.2.9 → autoinference_utils-0.2.11}/.gitignore +0 -0
- {autoinference_utils-0.2.9 → autoinference_utils-0.2.11}/README.md +0 -0
- {autoinference_utils-0.2.9 → autoinference_utils-0.2.11}/src/autoinference_utils/__init__.py +0 -0
- {autoinference_utils-0.2.9 → autoinference_utils-0.2.11}/src/autoinference_utils/router.py +0 -0
|
@@ -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
|
{autoinference_utils-0.2.9 → autoinference_utils-0.2.11}/src/autoinference_utils/endpoint.py
RENAMED
|
@@ -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
|
-
|
|
305
|
-
|
|
306
|
-
|
|
307
|
-
|
|
308
|
-
|
|
309
|
-
|
|
310
|
-
|
|
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
|
-
|
|
383
|
-
|
|
384
|
-
|
|
385
|
-
|
|
386
|
-
|
|
387
|
-
|
|
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() +
|
|
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
|
-
|
|
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])
|
|
File without changes
|
|
File without changes
|
{autoinference_utils-0.2.9 → autoinference_utils-0.2.11}/src/autoinference_utils/__init__.py
RENAMED
|
File without changes
|
|
File without changes
|