autoinference-utils 0.2.7__tar.gz → 0.2.8__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.7 → autoinference_utils-0.2.8}/PKG-INFO +1 -7
- autoinference_utils-0.2.8/README.md +9 -0
- {autoinference_utils-0.2.7 → autoinference_utils-0.2.8}/pyproject.toml +1 -1
- {autoinference_utils-0.2.7 → autoinference_utils-0.2.8}/src/autoinference_utils/endpoint.py +154 -117
- autoinference_utils-0.2.8/src/autoinference_utils/pd.py +315 -0
- autoinference_utils-0.2.7/README.md +0 -15
- autoinference_utils-0.2.7/src/autoinference_utils/model_path.py +0 -38
- autoinference_utils-0.2.7/src/autoinference_utils/tests/test_model_path.py +0 -84
- {autoinference_utils-0.2.7 → autoinference_utils-0.2.8}/.gitignore +0 -0
- {autoinference_utils-0.2.7 → autoinference_utils-0.2.8}/src/autoinference_utils/__init__.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.5
|
|
2
2
|
Name: autoinference-utils
|
|
3
|
-
Version: 0.2.
|
|
3
|
+
Version: 0.2.8
|
|
4
4
|
Summary: Shared endpoint abstractions for autoinference deployments
|
|
5
5
|
Requires-Python: >=3.10
|
|
6
6
|
Description-Content-Type: text/markdown
|
|
@@ -9,12 +9,6 @@ Description-Content-Type: text/markdown
|
|
|
9
9
|
|
|
10
10
|
Shared endpoint abstractions (`SGLangEndpoint`, `VLLMEndpoint`) for autoinference deployments.
|
|
11
11
|
|
|
12
|
-
For SGLang models whose custom Transformers code cannot load from symlinked
|
|
13
|
-
Hugging Face snapshots, set `AUTOINFERENCE_MATERIALIZE_MODEL_PATH=1`. The endpoint
|
|
14
|
-
copies local base/draft model code and metadata into unique temporary directories
|
|
15
|
-
and symlinks top-level safetensors weights. Repository IDs pass through unchanged.
|
|
16
|
-
Leave the setting unset to preserve the original paths.
|
|
17
|
-
|
|
18
12
|
## Publishing
|
|
19
13
|
|
|
20
14
|
Bump `version` in `pyproject.toml` and merge to `main` — CI publishes automatically via trusted publishing.
|
|
@@ -0,0 +1,9 @@
|
|
|
1
|
+
# autoinference-utils
|
|
2
|
+
|
|
3
|
+
Shared endpoint abstractions (`SGLangEndpoint`, `VLLMEndpoint`) for autoinference deployments.
|
|
4
|
+
|
|
5
|
+
## Publishing
|
|
6
|
+
|
|
7
|
+
Bump `version` in `pyproject.toml` and merge to `main` — CI publishes automatically via trusted publishing.
|
|
8
|
+
|
|
9
|
+
To publish manually: `cd autoinference_utils && uv build && uv publish`
|
|
@@ -23,18 +23,24 @@ import urllib.request
|
|
|
23
23
|
from abc import ABC
|
|
24
24
|
from typing import Any, Callable, Literal, Mapping, Optional, Sequence
|
|
25
25
|
|
|
26
|
-
from autoinference_utils.model_path import materialize_model_path
|
|
27
|
-
|
|
28
26
|
BENCH_MODE_ENV = "ENDPOINT_BENCH_MODE"
|
|
29
27
|
DUMMY_WEIGHTS_ENV = "ENDPOINT_DUMMY"
|
|
30
28
|
BENCH_MODE_PORT_OFFSET = 10000
|
|
31
29
|
|
|
32
|
-
MODAL_FLASH_REQUEST_UUID_HEADER
|
|
30
|
+
MODAL_FLASH_REQUEST_UUID_HEADER= "modal-flash-request-uuid"
|
|
33
31
|
MODAL_SESSION_ID_HEADER = "modal-session-id"
|
|
34
32
|
|
|
35
33
|
ENDPOINTS_REQUIRING_OAI_STREAM_COMPAT: tuple[str, ...] = ("TRTLLMEndpoint",)
|
|
36
34
|
|
|
37
35
|
|
|
36
|
+
def _url(host: str, port: int) -> str:
|
|
37
|
+
return f"http://{'[' + host + ']' if ':' in host else host}:{port}"
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def _local_base_url(host: str, port: int) -> str:
|
|
41
|
+
return _url({"0.0.0.0": "127.0.0.1", "::": "::1"}.get(host, host), port)
|
|
42
|
+
|
|
43
|
+
|
|
38
44
|
def server_arg_tokens(flag: str, value: str) -> list[str]:
|
|
39
45
|
"""Render one server arg as argv tokens.
|
|
40
46
|
|
|
@@ -57,7 +63,9 @@ def _flag_name(key: str) -> str:
|
|
|
57
63
|
return key.partition("=")[0]
|
|
58
64
|
|
|
59
65
|
|
|
60
|
-
def get_server_arg(
|
|
66
|
+
def get_server_arg(
|
|
67
|
+
server_args: Mapping[str, str], flag: str
|
|
68
|
+
) -> tuple[str, str] | None:
|
|
61
69
|
"""Return the key and value setting a flag under either spelling."""
|
|
62
70
|
for key, value in server_args.items():
|
|
63
71
|
if _flag_name(key) == flag:
|
|
@@ -150,18 +158,16 @@ class SGLangEndpoint(Endpoint):
|
|
|
150
158
|
sglang_port = (
|
|
151
159
|
worker_port + BENCH_MODE_PORT_OFFSET if self.bench_mode else worker_port
|
|
152
160
|
)
|
|
153
|
-
|
|
161
|
+
host_arg = get_server_arg(extra_server_args or {}, "--host")
|
|
162
|
+
host = (host_arg[1] or host_arg[0].partition("=")[2]) if host_arg else "0.0.0.0"
|
|
163
|
+
super().__init__(base_url=_local_base_url(host, sglang_port))
|
|
154
164
|
|
|
155
165
|
self.worker_port = sglang_port
|
|
156
|
-
self.model_path =
|
|
166
|
+
self.model_path = model_path
|
|
157
167
|
self.tp = tp
|
|
158
168
|
self.ep = ep
|
|
159
169
|
self.dp = dp
|
|
160
|
-
self.speculative_model_path =
|
|
161
|
-
materialize_model_path(speculative_model_path)
|
|
162
|
-
if speculative_model_path is not None
|
|
163
|
-
else None
|
|
164
|
-
)
|
|
170
|
+
self.speculative_model_path = speculative_model_path
|
|
165
171
|
self.load_format = load_format
|
|
166
172
|
self.nnodes = nnodes
|
|
167
173
|
self.node_rank = node_rank
|
|
@@ -170,7 +176,9 @@ class SGLangEndpoint(Endpoint):
|
|
|
170
176
|
self.disaggregation_mode = disaggregation_mode
|
|
171
177
|
self.prefill_bootstrap_port = prefill_bootstrap_port
|
|
172
178
|
self.launcher_module = launcher_module
|
|
173
|
-
self.extra_server_args =
|
|
179
|
+
self.extra_server_args = (
|
|
180
|
+
dict(extra_server_args) if extra_server_args else {}
|
|
181
|
+
)
|
|
174
182
|
self.health_timeout = health_timeout
|
|
175
183
|
self.health_poll_interval = health_poll_interval
|
|
176
184
|
self.health_request_timeout = health_request_timeout
|
|
@@ -185,16 +193,17 @@ class SGLangEndpoint(Endpoint):
|
|
|
185
193
|
self._bench_server: Optional[http.server.ThreadingHTTPServer] = None
|
|
186
194
|
|
|
187
195
|
if self.disaggregation_mode not in (None, "prefill", "decode"):
|
|
188
|
-
raise ValueError(
|
|
196
|
+
raise ValueError(
|
|
197
|
+
"disaggregation_mode must be None, 'prefill', or 'decode'"
|
|
198
|
+
)
|
|
189
199
|
|
|
190
200
|
if os.environ.get(DUMMY_WEIGHTS_ENV) == "1":
|
|
191
|
-
if
|
|
192
|
-
self.
|
|
201
|
+
if (
|
|
202
|
+
self.load_format is None
|
|
203
|
+
and not has_server_arg(self.extra_server_args, "--load-format")
|
|
193
204
|
):
|
|
194
205
|
self.load_format = "dummy"
|
|
195
|
-
if not has_server_arg(
|
|
196
|
-
self.extra_server_args, "--model-loader-extra-config"
|
|
197
|
-
):
|
|
206
|
+
if not has_server_arg(self.extra_server_args, "--model-loader-extra-config"):
|
|
198
207
|
self.extra_server_args["--model-loader-extra-config"] = "{}"
|
|
199
208
|
|
|
200
209
|
# Request logging enabled, log only metadata by default
|
|
@@ -205,9 +214,7 @@ class SGLangEndpoint(Endpoint):
|
|
|
205
214
|
# Set request logging level
|
|
206
215
|
level = get_server_arg(self.extra_server_args, "--log-requests-level")
|
|
207
216
|
if level is None:
|
|
208
|
-
self.extra_server_args["--log-requests-level"] = str(
|
|
209
|
-
self.log_requests_level
|
|
210
|
-
)
|
|
217
|
+
self.extra_server_args["--log-requests-level"] = str(self.log_requests_level)
|
|
211
218
|
else:
|
|
212
219
|
key, value = level
|
|
213
220
|
print(
|
|
@@ -229,21 +236,19 @@ class SGLangEndpoint(Endpoint):
|
|
|
229
236
|
self.log_request_headers = ",".join(sglang_log_request_headers)
|
|
230
237
|
os.environ["SGLANG_LOG_REQUEST_HEADERS"] = self.log_request_headers
|
|
231
238
|
|
|
239
|
+
|
|
232
240
|
def _build_cmd(self) -> list[str]:
|
|
233
241
|
cmd = [
|
|
234
|
-
"python",
|
|
235
|
-
"
|
|
236
|
-
self.
|
|
237
|
-
"--
|
|
238
|
-
"0.0.0.0",
|
|
239
|
-
"--port",
|
|
240
|
-
str(self.worker_port),
|
|
241
|
-
"--model-path",
|
|
242
|
-
self.model_path,
|
|
242
|
+
"python", "-m", self.launcher_module,
|
|
243
|
+
"--host", "0.0.0.0",
|
|
244
|
+
"--port", str(self.worker_port),
|
|
245
|
+
"--model-path", self.model_path,
|
|
243
246
|
]
|
|
244
247
|
|
|
245
248
|
if self.speculative_model_path is not None:
|
|
246
|
-
cmd.extend(
|
|
249
|
+
cmd.extend(
|
|
250
|
+
["--speculative-draft-model-path", self.speculative_model_path]
|
|
251
|
+
)
|
|
247
252
|
if self.load_format is not None:
|
|
248
253
|
cmd.extend(["--load-format", self.load_format])
|
|
249
254
|
if self.tp is not None:
|
|
@@ -268,25 +273,21 @@ class SGLangEndpoint(Endpoint):
|
|
|
268
273
|
raise ValueError("dist_init_host is required when nnodes > 1")
|
|
269
274
|
cmd.extend(
|
|
270
275
|
[
|
|
271
|
-
"--nnodes",
|
|
272
|
-
str(self.
|
|
273
|
-
"--node-rank",
|
|
274
|
-
str(self.node_rank),
|
|
276
|
+
"--nnodes", str(self.nnodes),
|
|
277
|
+
"--node-rank", str(self.node_rank),
|
|
275
278
|
"--dist-init-addr",
|
|
276
279
|
f"{self.dist_init_host}:{self.dist_init_port}",
|
|
277
280
|
]
|
|
278
281
|
)
|
|
279
282
|
|
|
280
|
-
merged = _merge_server_args(
|
|
281
|
-
self.DEFAULT_OPERATIONAL_ARGS, self.extra_server_args
|
|
282
|
-
)
|
|
283
|
+
merged = _merge_server_args(self.DEFAULT_OPERATIONAL_ARGS, self.extra_server_args)
|
|
283
284
|
for key, value in merged.items():
|
|
284
285
|
cmd.extend(server_arg_tokens(key, value))
|
|
285
286
|
|
|
286
287
|
return cmd
|
|
287
288
|
|
|
288
289
|
def health_check(self) -> str | None:
|
|
289
|
-
url = f"
|
|
290
|
+
url = f"{self.base_url}/health"
|
|
290
291
|
return _health_check(
|
|
291
292
|
url,
|
|
292
293
|
request_timeout=self.health_request_timeout,
|
|
@@ -300,6 +301,7 @@ class SGLangEndpoint(Endpoint):
|
|
|
300
301
|
wait_ready(
|
|
301
302
|
self._proc,
|
|
302
303
|
port=self.worker_port,
|
|
304
|
+
base_url=self.base_url,
|
|
303
305
|
timeout=self.health_timeout,
|
|
304
306
|
poll_interval=self.health_poll_interval,
|
|
305
307
|
request_timeout=self.health_request_timeout,
|
|
@@ -308,6 +310,7 @@ class SGLangEndpoint(Endpoint):
|
|
|
308
310
|
self._bench_server = start_bench_proxy(
|
|
309
311
|
listen_port=self.listen_port,
|
|
310
312
|
upstream_port=self.worker_port,
|
|
313
|
+
upstream_base_url=self.base_url,
|
|
311
314
|
)
|
|
312
315
|
|
|
313
316
|
def stop(self):
|
|
@@ -342,7 +345,9 @@ class VLLMEndpoint(Endpoint):
|
|
|
342
345
|
super().__init__(base_url=f"http://localhost:{vllm_port}")
|
|
343
346
|
self.model = model
|
|
344
347
|
self.worker_port = vllm_port
|
|
345
|
-
self.extra_server_args =
|
|
348
|
+
self.extra_server_args = (
|
|
349
|
+
dict(extra_server_args) if extra_server_args else {}
|
|
350
|
+
)
|
|
346
351
|
self.health_timeout = health_timeout
|
|
347
352
|
self.health_poll_interval = health_poll_interval
|
|
348
353
|
self.health_request_timeout = health_request_timeout
|
|
@@ -355,15 +360,10 @@ class VLLMEndpoint(Endpoint):
|
|
|
355
360
|
|
|
356
361
|
def _build_cmd(self) -> list[str]:
|
|
357
362
|
cmd = [
|
|
358
|
-
"python",
|
|
359
|
-
"
|
|
360
|
-
"
|
|
361
|
-
"--
|
|
362
|
-
"0.0.0.0",
|
|
363
|
-
"--port",
|
|
364
|
-
str(self.worker_port),
|
|
365
|
-
"--model",
|
|
366
|
-
self.model,
|
|
363
|
+
"python", "-m", "vllm.entrypoints.openai.api_server",
|
|
364
|
+
"--host", "0.0.0.0",
|
|
365
|
+
"--port", str(self.worker_port),
|
|
366
|
+
"--model", self.model,
|
|
367
367
|
]
|
|
368
368
|
for key, value in self.extra_server_args.items():
|
|
369
369
|
cmd.extend(server_arg_tokens(key, value))
|
|
@@ -417,7 +417,9 @@ class TRTLLMEndpoint(Endpoint):
|
|
|
417
417
|
super().__init__(base_url=f"http://localhost:{worker_port}")
|
|
418
418
|
self.model = model
|
|
419
419
|
self.worker_port = worker_port
|
|
420
|
-
self.extra_server_args =
|
|
420
|
+
self.extra_server_args = (
|
|
421
|
+
dict(extra_server_args) if extra_server_args else {}
|
|
422
|
+
)
|
|
421
423
|
self.health_timeout = health_timeout
|
|
422
424
|
self.health_poll_interval = health_poll_interval
|
|
423
425
|
self.health_request_timeout = health_request_timeout
|
|
@@ -427,10 +429,8 @@ class TRTLLMEndpoint(Endpoint):
|
|
|
427
429
|
cmd = [
|
|
428
430
|
"trtllm-serve",
|
|
429
431
|
self.model,
|
|
430
|
-
"--host",
|
|
431
|
-
"
|
|
432
|
-
"--port",
|
|
433
|
-
str(self.worker_port),
|
|
432
|
+
"--host", "0.0.0.0",
|
|
433
|
+
"--port", str(self.worker_port),
|
|
434
434
|
]
|
|
435
435
|
for key, value in self.extra_server_args.items():
|
|
436
436
|
cmd.extend(server_arg_tokens(key, value))
|
|
@@ -474,13 +474,16 @@ class RouterEndpoint(Endpoint):
|
|
|
474
474
|
api_key: Optional[str] = None,
|
|
475
475
|
health_timeout: float = 10 * 60,
|
|
476
476
|
health_poll_interval: float = 5.0,
|
|
477
|
+
health_path: str = "/health",
|
|
478
|
+
host: str = "0.0.0.0",
|
|
477
479
|
):
|
|
478
480
|
self.bench_mode = os.environ.get(BENCH_MODE_ENV) == "1"
|
|
479
481
|
self.listen_port = router_port
|
|
480
482
|
actual_router_port = (
|
|
481
483
|
router_port + BENCH_MODE_PORT_OFFSET if self.bench_mode else router_port
|
|
482
484
|
)
|
|
483
|
-
|
|
485
|
+
self.host = host
|
|
486
|
+
super().__init__(base_url=_local_base_url(host, actual_router_port))
|
|
484
487
|
self.pd_config = list(pd_config)
|
|
485
488
|
self.worker_port = (
|
|
486
489
|
worker_port + BENCH_MODE_PORT_OFFSET if self.bench_mode else worker_port
|
|
@@ -490,35 +493,24 @@ class RouterEndpoint(Endpoint):
|
|
|
490
493
|
self.api_key = api_key
|
|
491
494
|
self.health_timeout = health_timeout
|
|
492
495
|
self.health_poll_interval = health_poll_interval
|
|
496
|
+
self.health_path = health_path
|
|
493
497
|
self._proc: Optional[subprocess.Popen] = None
|
|
494
498
|
self._bench_server: Optional[http.server.ThreadingHTTPServer] = None
|
|
495
499
|
|
|
496
500
|
def _build_cmd(self) -> list[str]:
|
|
497
501
|
cmd = [
|
|
498
|
-
"python",
|
|
499
|
-
"
|
|
500
|
-
"
|
|
501
|
-
"--
|
|
502
|
-
"
|
|
503
|
-
"--
|
|
504
|
-
|
|
505
|
-
"--
|
|
506
|
-
"
|
|
507
|
-
"--
|
|
508
|
-
"round_robin",
|
|
509
|
-
"--max-concurrent-requests",
|
|
510
|
-
"128",
|
|
511
|
-
"--rate-limit-tokens-per-second",
|
|
512
|
-
"0",
|
|
513
|
-
"--queue-size",
|
|
514
|
-
"0",
|
|
515
|
-
"--health-check-timeout-secs",
|
|
516
|
-
"600",
|
|
517
|
-
"--log-level",
|
|
518
|
-
"info",
|
|
502
|
+
"python", "-m", "sglang_router.launch_router",
|
|
503
|
+
"--host", f"[{self.host}]" if ":" in self.host else self.host,
|
|
504
|
+
"--port", str(self.router_port),
|
|
505
|
+
"--prefill-policy", "cache_aware",
|
|
506
|
+
"--decode-policy", "round_robin",
|
|
507
|
+
"--max-concurrent-requests", "128",
|
|
508
|
+
"--rate-limit-tokens-per-second", "0",
|
|
509
|
+
"--queue-size", "0",
|
|
510
|
+
"--health-check-timeout-secs", "600",
|
|
511
|
+
"--log-level", "info",
|
|
519
512
|
"--disable-circuit-breaker",
|
|
520
|
-
"--request-timeout-secs",
|
|
521
|
-
"3600",
|
|
513
|
+
"--request-timeout-secs", "3600",
|
|
522
514
|
]
|
|
523
515
|
|
|
524
516
|
if self.api_key is not None:
|
|
@@ -528,9 +520,11 @@ class RouterEndpoint(Endpoint):
|
|
|
528
520
|
cmd.append("--pd-disaggregation")
|
|
529
521
|
|
|
530
522
|
for role, node_ip in self.pd_config:
|
|
531
|
-
node_url =
|
|
523
|
+
node_url = _url(node_ip, self.worker_port)
|
|
532
524
|
if role == "prefill":
|
|
533
|
-
cmd.extend(
|
|
525
|
+
cmd.extend(
|
|
526
|
+
["--prefill", node_url, str(self.prefill_bootstrap_port)]
|
|
527
|
+
)
|
|
534
528
|
elif role == "decode":
|
|
535
529
|
cmd.extend(["--decode", node_url])
|
|
536
530
|
elif role == "worker":
|
|
@@ -543,7 +537,7 @@ class RouterEndpoint(Endpoint):
|
|
|
543
537
|
def start(self):
|
|
544
538
|
for _, node_ip in self.pd_config:
|
|
545
539
|
_wait_ready_url(
|
|
546
|
-
f"
|
|
540
|
+
f"{_url(node_ip, self.worker_port)}/health",
|
|
547
541
|
timeout=self.health_timeout,
|
|
548
542
|
poll_interval=self.health_poll_interval,
|
|
549
543
|
)
|
|
@@ -551,17 +545,25 @@ class RouterEndpoint(Endpoint):
|
|
|
551
545
|
cmd = self._build_cmd()
|
|
552
546
|
print(f"[router] starting: {shlex.join(cmd)}")
|
|
553
547
|
self._proc = subprocess.Popen(cmd)
|
|
554
|
-
|
|
555
|
-
f"http://localhost:{self.router_port}/health",
|
|
556
|
-
timeout=self.health_timeout,
|
|
557
|
-
poll_interval=self.health_poll_interval,
|
|
558
|
-
)
|
|
548
|
+
self.wait_ready()
|
|
559
549
|
if self.bench_mode:
|
|
560
550
|
self._bench_server = start_bench_proxy(
|
|
561
551
|
listen_port=self.listen_port,
|
|
562
552
|
upstream_port=self.router_port,
|
|
553
|
+
upstream_base_url=self.base_url,
|
|
563
554
|
)
|
|
564
555
|
|
|
556
|
+
def wait_ready(self) -> None:
|
|
557
|
+
assert self._proc is not None
|
|
558
|
+
wait_ready(
|
|
559
|
+
self._proc,
|
|
560
|
+
port=self.router_port,
|
|
561
|
+
base_url=self.base_url,
|
|
562
|
+
health_path=self.health_path,
|
|
563
|
+
timeout=self.health_timeout,
|
|
564
|
+
poll_interval=self.health_poll_interval,
|
|
565
|
+
)
|
|
566
|
+
|
|
565
567
|
def stop(self):
|
|
566
568
|
if self._bench_server is not None:
|
|
567
569
|
self._bench_server.shutdown()
|
|
@@ -574,28 +576,34 @@ class RouterEndpoint(Endpoint):
|
|
|
574
576
|
# and this proxy listens on worker_port. POST /bench shells out to a benchmark
|
|
575
577
|
# task via run_bench(); every other path is forwarded to the upstream server.
|
|
576
578
|
|
|
577
|
-
|
|
578
579
|
def start_bench_proxy(
|
|
579
580
|
*,
|
|
580
581
|
listen_port: int,
|
|
581
582
|
upstream_port: int,
|
|
583
|
+
upstream_base_url: str | None = None,
|
|
582
584
|
) -> http.server.ThreadingHTTPServer:
|
|
583
585
|
handler_cls = type(
|
|
584
586
|
"BenchProxyHandler",
|
|
585
587
|
(_BenchProxyHandler,),
|
|
586
|
-
{"upstream_port": upstream_port},
|
|
588
|
+
{"upstream_port": upstream_port, "upstream_base_url": upstream_base_url},
|
|
589
|
+
)
|
|
590
|
+
server = http.server.ThreadingHTTPServer(
|
|
591
|
+
("0.0.0.0", listen_port), handler_cls
|
|
587
592
|
)
|
|
588
|
-
server = http.server.ThreadingHTTPServer(("0.0.0.0", listen_port), handler_cls)
|
|
589
593
|
thread = threading.Thread(
|
|
590
594
|
target=server.serve_forever, daemon=True, name="bench-proxy"
|
|
591
595
|
)
|
|
592
596
|
thread.start()
|
|
593
|
-
print(
|
|
597
|
+
print(
|
|
598
|
+
f"[bench-proxy] listening on :{listen_port}, "
|
|
599
|
+
f"forwarding to :{upstream_port}"
|
|
600
|
+
)
|
|
594
601
|
return server
|
|
595
602
|
|
|
596
603
|
|
|
597
604
|
class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
|
|
598
605
|
upstream_port: int = 8000
|
|
606
|
+
upstream_base_url: str | None = None
|
|
599
607
|
|
|
600
608
|
def log_message(self, format, *args):
|
|
601
609
|
return
|
|
@@ -619,24 +627,35 @@ class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
|
|
|
619
627
|
try:
|
|
620
628
|
payload = self._read_json_body()
|
|
621
629
|
except (ValueError, json.JSONDecodeError) as exc:
|
|
622
|
-
return self._send_json(
|
|
630
|
+
return self._send_json(
|
|
631
|
+
400, {"ok": False, "error": f"bad body: {exc}"}
|
|
632
|
+
)
|
|
623
633
|
|
|
624
634
|
benchmark = payload.get("benchmark") or ""
|
|
625
635
|
args = payload.get("args") or []
|
|
626
|
-
target =
|
|
636
|
+
target = (
|
|
637
|
+
payload.get("target")
|
|
638
|
+
or self.upstream_base_url
|
|
639
|
+
or f"http://localhost:{self.upstream_port}"
|
|
640
|
+
)
|
|
627
641
|
output_dir = payload.get("output_dir") or "/tmp/bench-output"
|
|
628
642
|
|
|
629
643
|
if not benchmark:
|
|
630
|
-
return self._send_json(
|
|
644
|
+
return self._send_json(
|
|
645
|
+
400, {"ok": False, "error": "benchmark is required"}
|
|
646
|
+
)
|
|
631
647
|
if not isinstance(args, list):
|
|
632
|
-
return self._send_json(
|
|
648
|
+
return self._send_json(
|
|
649
|
+
400, {"ok": False, "error": "args must be a list"}
|
|
650
|
+
)
|
|
633
651
|
|
|
634
652
|
rc, body = run_bench(benchmark, list(args), target, output_dir)
|
|
635
653
|
status = 200 if rc == 0 else 500
|
|
636
654
|
self._send_raw(status, "application/json", body)
|
|
637
655
|
|
|
638
656
|
def _proxy(self, method: str):
|
|
639
|
-
|
|
657
|
+
base_url = self.upstream_base_url or f"http://localhost:{self.upstream_port}"
|
|
658
|
+
url = f"{base_url}{self.path}"
|
|
640
659
|
length = int(self.headers.get("Content-Length", 0))
|
|
641
660
|
body = self.rfile.read(length) if length else None
|
|
642
661
|
forward_headers = {
|
|
@@ -680,7 +699,9 @@ class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
|
|
|
680
699
|
return json.loads(raw)
|
|
681
700
|
|
|
682
701
|
def _send_json(self, status: int, obj: Mapping[str, Any]):
|
|
683
|
-
self._send_raw(
|
|
702
|
+
self._send_raw(
|
|
703
|
+
status, "application/json", json.dumps(obj).encode()
|
|
704
|
+
)
|
|
684
705
|
|
|
685
706
|
def _send_raw(self, status: int, content_type: str, body: bytes):
|
|
686
707
|
self.send_response(status)
|
|
@@ -705,21 +726,18 @@ def run_bench(
|
|
|
705
726
|
return 500, json.dumps(
|
|
706
727
|
{
|
|
707
728
|
"ok": False,
|
|
708
|
-
"error": (
|
|
729
|
+
"error": (
|
|
730
|
+
"autoinference.benchmarks not importable in container"
|
|
731
|
+
),
|
|
709
732
|
"volume_path": output_dir,
|
|
710
733
|
}
|
|
711
734
|
).encode()
|
|
712
735
|
|
|
713
736
|
cmd = [
|
|
714
|
-
sys.executable,
|
|
715
|
-
"-
|
|
716
|
-
"
|
|
717
|
-
|
|
718
|
-
search_root,
|
|
719
|
-
"-c",
|
|
720
|
-
"tasks",
|
|
721
|
-
benchmark,
|
|
722
|
-
*[str(a) for a in args],
|
|
737
|
+
sys.executable, "-m", "invoke",
|
|
738
|
+
"--search-root", search_root,
|
|
739
|
+
"-c", "tasks",
|
|
740
|
+
benchmark, *[str(a) for a in args],
|
|
723
741
|
]
|
|
724
742
|
env = {
|
|
725
743
|
**os.environ,
|
|
@@ -759,17 +777,20 @@ def wait_ready(
|
|
|
759
777
|
port: int,
|
|
760
778
|
timeout: float,
|
|
761
779
|
health_path: str = "/health",
|
|
780
|
+
base_url: str | None = None,
|
|
762
781
|
poll_interval: float = 5.0,
|
|
763
782
|
request_timeout: float = 5.0,
|
|
764
783
|
) -> None:
|
|
765
784
|
"""Poll an HTTP health endpoint until ready, raising if the process dies."""
|
|
766
785
|
deadline = time.time() + timeout
|
|
767
|
-
url = f"
|
|
786
|
+
url = f"{base_url or _local_base_url('127.0.0.1', port)}{health_path}"
|
|
768
787
|
last_error = "no response yet"
|
|
769
788
|
|
|
770
789
|
while time.time() < deadline:
|
|
771
790
|
try:
|
|
772
|
-
error = _health_check(
|
|
791
|
+
error = _health_check(
|
|
792
|
+
url, request_timeout=request_timeout, process=process
|
|
793
|
+
)
|
|
773
794
|
except subprocess.CalledProcessError as exc:
|
|
774
795
|
print(
|
|
775
796
|
f"[endpoint] !!! server process exited with code "
|
|
@@ -783,17 +804,20 @@ def wait_ready(
|
|
|
783
804
|
time.sleep(poll_interval)
|
|
784
805
|
|
|
785
806
|
print(
|
|
786
|
-
f"[endpoint] !!! health check timed out after {timeout}s;
|
|
807
|
+
f"[endpoint] !!! health check timed out after {timeout}s; "
|
|
808
|
+
f"last={last_error}",
|
|
787
809
|
flush=True,
|
|
788
810
|
)
|
|
789
811
|
raise TimeoutError(
|
|
790
|
-
f"Health check timed out after {timeout}s for {url}.
|
|
812
|
+
f"Health check timed out after {timeout}s for {url}. "
|
|
813
|
+
f"Last error: {last_error}"
|
|
791
814
|
)
|
|
792
815
|
|
|
793
816
|
|
|
794
817
|
def warmup_chat_completions(
|
|
795
818
|
*,
|
|
796
819
|
port: int,
|
|
820
|
+
base_url: str | None = None,
|
|
797
821
|
payload: Mapping[str, Any],
|
|
798
822
|
headers: Mapping[str, str] | None = None,
|
|
799
823
|
successful_requests: int = 3,
|
|
@@ -802,7 +826,7 @@ def warmup_chat_completions(
|
|
|
802
826
|
retry_delay: float = 1.0,
|
|
803
827
|
) -> None:
|
|
804
828
|
"""Warm the OpenAI chat completions endpoint with strict retries."""
|
|
805
|
-
url = f"
|
|
829
|
+
url = f"{base_url or _local_base_url('127.0.0.1', port)}/v1/chat/completions"
|
|
806
830
|
request_headers = {"Content-Type": "application/json"}
|
|
807
831
|
if headers:
|
|
808
832
|
request_headers.update(headers)
|
|
@@ -815,7 +839,9 @@ def warmup_chat_completions(
|
|
|
815
839
|
request_timeout=request_timeout,
|
|
816
840
|
max_attempts=max_attempts_per_request,
|
|
817
841
|
retry_delay=retry_delay,
|
|
818
|
-
description=(
|
|
842
|
+
description=(
|
|
843
|
+
f"warmup request {request_idx + 1}/{successful_requests}"
|
|
844
|
+
),
|
|
819
845
|
)
|
|
820
846
|
|
|
821
847
|
|
|
@@ -837,7 +863,9 @@ def validate_embeddings_endpoint(
|
|
|
837
863
|
1 if all(isinstance(item, int) for item in inputs) else len(inputs)
|
|
838
864
|
)
|
|
839
865
|
else:
|
|
840
|
-
raise ValueError(
|
|
866
|
+
raise ValueError(
|
|
867
|
+
"embedding validation payload must contain non-empty input"
|
|
868
|
+
)
|
|
841
869
|
|
|
842
870
|
request_headers = {"Content-Type": "application/json"}
|
|
843
871
|
if headers:
|
|
@@ -900,7 +928,10 @@ def start_heartbeat_thread(
|
|
|
900
928
|
try:
|
|
901
929
|
error = health_check_fn()
|
|
902
930
|
except subprocess.CalledProcessError as exc:
|
|
903
|
-
print(
|
|
931
|
+
print(
|
|
932
|
+
f"[heartbeat] server process exited with code "
|
|
933
|
+
f"{exc.returncode}"
|
|
934
|
+
)
|
|
904
935
|
on_failure()
|
|
905
936
|
return
|
|
906
937
|
if error is None:
|
|
@@ -912,7 +943,10 @@ def start_heartbeat_thread(
|
|
|
912
943
|
f"({consecutive_failures}/{max_consecutive_failures})"
|
|
913
944
|
)
|
|
914
945
|
if consecutive_failures >= max_consecutive_failures:
|
|
915
|
-
print(
|
|
946
|
+
print(
|
|
947
|
+
"[heartbeat] sustained health-check failure, "
|
|
948
|
+
"invoking on_failure"
|
|
949
|
+
)
|
|
916
950
|
on_failure()
|
|
917
951
|
return
|
|
918
952
|
|
|
@@ -997,7 +1031,8 @@ def _wait_ready_url(
|
|
|
997
1031
|
last_error = error
|
|
998
1032
|
time.sleep(poll_interval)
|
|
999
1033
|
raise TimeoutError(
|
|
1000
|
-
f"Timed out after {timeout}s waiting for {url}.
|
|
1034
|
+
f"Timed out after {timeout}s waiting for {url}. "
|
|
1035
|
+
f"Last error: {last_error}"
|
|
1001
1036
|
)
|
|
1002
1037
|
|
|
1003
1038
|
|
|
@@ -1015,7 +1050,9 @@ def _post_json(
|
|
|
1015
1050
|
timeout: float | None = None,
|
|
1016
1051
|
) -> bytes:
|
|
1017
1052
|
body = json.dumps(payload).encode("utf-8")
|
|
1018
|
-
req = urllib.request.Request(
|
|
1053
|
+
req = urllib.request.Request(
|
|
1054
|
+
url, data=body, headers=dict(headers), method="POST"
|
|
1055
|
+
)
|
|
1019
1056
|
with urllib.request.urlopen(req, timeout=timeout) as resp:
|
|
1020
1057
|
return resp.read()
|
|
1021
1058
|
|
|
@@ -0,0 +1,315 @@
|
|
|
1
|
+
"""SGLang p/d disaggregated serving in a multi-container static cluster."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import asyncio
|
|
6
|
+
import math
|
|
7
|
+
import os
|
|
8
|
+
import subprocess
|
|
9
|
+
import threading
|
|
10
|
+
import time
|
|
11
|
+
from collections.abc import Callable, Mapping
|
|
12
|
+
|
|
13
|
+
from .endpoint import (
|
|
14
|
+
Endpoint,
|
|
15
|
+
RouterEndpoint,
|
|
16
|
+
SGLangEndpoint,
|
|
17
|
+
_flag_name,
|
|
18
|
+
_health_check,
|
|
19
|
+
_local_base_url,
|
|
20
|
+
_merge_server_args,
|
|
21
|
+
_url,
|
|
22
|
+
)
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def _configure_fabric() -> None:
|
|
26
|
+
if os.environ.get("FI_PROVIDER") == "efa":
|
|
27
|
+
os.environ["SGLANG_DISAGGREGATION_NIXL_BACKEND"] = "LIBFABRIC"
|
|
28
|
+
elif hca := os.environ.get("NCCL_IB_HCA", ""):
|
|
29
|
+
rails = [name for name in hca.lstrip("=").split(",") if name]
|
|
30
|
+
os.environ["UCX_NET_DEVICES"] = ",".join(f"{name}:1" for name in rails)
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
async def _stop_container() -> None:
|
|
34
|
+
from modal.client import _Client
|
|
35
|
+
from modal_proto import api_pb2
|
|
36
|
+
|
|
37
|
+
async def stop():
|
|
38
|
+
client = await _Client.from_env()
|
|
39
|
+
await client.stub.ContainerStop(
|
|
40
|
+
api_pb2.ContainerStopRequest(task_id=os.environ["MODAL_TASK_ID"])
|
|
41
|
+
)
|
|
42
|
+
|
|
43
|
+
await asyncio.wait_for(stop(), timeout=3)
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
def _terminate_container() -> None:
|
|
47
|
+
from modal._utils.async_utils import synchronize_api
|
|
48
|
+
|
|
49
|
+
try:
|
|
50
|
+
synchronize_api(_stop_container)()
|
|
51
|
+
finally:
|
|
52
|
+
os._exit(1)
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
class _PDRouter(RouterEndpoint):
|
|
56
|
+
def __init__(
|
|
57
|
+
self,
|
|
58
|
+
*,
|
|
59
|
+
max_concurrent_requests,
|
|
60
|
+
drain_timeout,
|
|
61
|
+
health_interval,
|
|
62
|
+
health_request_timeout,
|
|
63
|
+
health_failure_threshold,
|
|
64
|
+
**kwargs,
|
|
65
|
+
):
|
|
66
|
+
super().__init__(health_path="/readiness", host="::", **kwargs)
|
|
67
|
+
self.max_concurrent_requests = max_concurrent_requests
|
|
68
|
+
self.drain_timeout = drain_timeout
|
|
69
|
+
self.health_interval = health_interval
|
|
70
|
+
self.health_request_timeout = health_request_timeout
|
|
71
|
+
self.health_failure_threshold = health_failure_threshold
|
|
72
|
+
|
|
73
|
+
def _build_cmd(self) -> list[str]:
|
|
74
|
+
cmd = [
|
|
75
|
+
"python",
|
|
76
|
+
"-m",
|
|
77
|
+
"sglang_router.launch_router",
|
|
78
|
+
"--host",
|
|
79
|
+
"[::]",
|
|
80
|
+
"--port",
|
|
81
|
+
str(self.router_port),
|
|
82
|
+
"--pd-disaggregation",
|
|
83
|
+
]
|
|
84
|
+
for role, host in self.pd_config:
|
|
85
|
+
cmd.extend([f"--{role}", _url(host, self.worker_port)])
|
|
86
|
+
if role == "prefill":
|
|
87
|
+
cmd.append(str(self.prefill_bootstrap_port))
|
|
88
|
+
for flag, value in {
|
|
89
|
+
"--prefill-policy": "cache_aware",
|
|
90
|
+
"--decode-policy": "round_robin",
|
|
91
|
+
"--max-concurrent-requests": self.max_concurrent_requests,
|
|
92
|
+
"--queue-size": 0,
|
|
93
|
+
"--rate-limit-tokens-per-second": 0,
|
|
94
|
+
"--health-check-interval-secs": self.health_interval,
|
|
95
|
+
"--health-check-timeout-secs": self.health_request_timeout,
|
|
96
|
+
"--health-failure-threshold": self.health_failure_threshold,
|
|
97
|
+
"--shutdown-grace-period-secs": self.drain_timeout,
|
|
98
|
+
"--request-timeout-secs": 1800,
|
|
99
|
+
"--log-level": "info",
|
|
100
|
+
}.items():
|
|
101
|
+
cmd.extend([flag, str(value)])
|
|
102
|
+
if self.api_key is not None:
|
|
103
|
+
cmd.extend(["--api-key", self.api_key])
|
|
104
|
+
return cmd
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
class PDEndpoint(Endpoint):
|
|
108
|
+
"""Own a local SGLang engine and the rank-zero router for P/D serving.
|
|
109
|
+
|
|
110
|
+
Call start/stop from Modal enter/exit hooks. Model-specific arguments and
|
|
111
|
+
warmup stay in the recipe. The image must support NIXL and graceful shutdown.
|
|
112
|
+
The ratio is (prefill, decode) counts; prefill ranks precede decode ranks.
|
|
113
|
+
"""
|
|
114
|
+
|
|
115
|
+
def __init__(
|
|
116
|
+
self,
|
|
117
|
+
engine: SGLangEndpoint,
|
|
118
|
+
*,
|
|
119
|
+
ratio: tuple[int, int],
|
|
120
|
+
prefill_args: Mapping[str, str] | None = None,
|
|
121
|
+
decode_args: Mapping[str, str] | None = None,
|
|
122
|
+
router_port: int = 9000,
|
|
123
|
+
max_concurrent_requests: int = 96,
|
|
124
|
+
api_key: str | None = None,
|
|
125
|
+
drain_timeout: int = 120,
|
|
126
|
+
health_interval: int = 10,
|
|
127
|
+
health_request_timeout: int = 10,
|
|
128
|
+
router_health_request_timeout: int = 5,
|
|
129
|
+
health_failure_threshold: int = 3,
|
|
130
|
+
health_failure_timeout: float = 30,
|
|
131
|
+
):
|
|
132
|
+
from modal.experimental import get_cluster_info
|
|
133
|
+
|
|
134
|
+
if not isinstance(engine, SGLangEndpoint):
|
|
135
|
+
raise TypeError("PDEndpoint requires a SGLangEndpoint")
|
|
136
|
+
if engine._proc is not None or engine.bench_mode:
|
|
137
|
+
raise ValueError("PDEndpoint requires an unstarted, non-benchmark engine")
|
|
138
|
+
if (
|
|
139
|
+
min(
|
|
140
|
+
health_interval,
|
|
141
|
+
health_request_timeout,
|
|
142
|
+
router_health_request_timeout,
|
|
143
|
+
health_failure_threshold,
|
|
144
|
+
)
|
|
145
|
+
<= 0
|
|
146
|
+
or drain_timeout < 0
|
|
147
|
+
or not math.isfinite(health_failure_timeout)
|
|
148
|
+
or health_failure_timeout <= 0
|
|
149
|
+
):
|
|
150
|
+
raise ValueError(
|
|
151
|
+
"health settings must be positive and drain_timeout nonnegative"
|
|
152
|
+
)
|
|
153
|
+
if len(ratio) != 2 or any(
|
|
154
|
+
type(count) is not int or count < 1 for count in ratio
|
|
155
|
+
):
|
|
156
|
+
raise ValueError(
|
|
157
|
+
"ratio must contain positive integer prefill and decode counts"
|
|
158
|
+
)
|
|
159
|
+
roles = ("prefill",) * ratio[0] + ("decode",) * ratio[1]
|
|
160
|
+
cluster = get_cluster_info()
|
|
161
|
+
hosts = cluster.container_ips
|
|
162
|
+
if len(hosts) != len(roles) or not 0 <= cluster.rank < len(roles):
|
|
163
|
+
raise ValueError(
|
|
164
|
+
f"P/D ratio {ratio} requires a {len(roles)}-container Modal cluster"
|
|
165
|
+
)
|
|
166
|
+
self._host_ip = hosts[cluster.rank]
|
|
167
|
+
self.role = roles[cluster.rank]
|
|
168
|
+
super().__init__(_url(hosts[0], router_port))
|
|
169
|
+
self.engine = engine
|
|
170
|
+
engine.disaggregation_mode = self.role
|
|
171
|
+
engine.health_request_timeout = health_request_timeout
|
|
172
|
+
engine.extra_server_args = _merge_server_args(
|
|
173
|
+
engine.extra_server_args,
|
|
174
|
+
(prefill_args if self.role == "prefill" else decode_args) or {},
|
|
175
|
+
)
|
|
176
|
+
engine.extra_server_args = _merge_server_args(
|
|
177
|
+
engine.extra_server_args,
|
|
178
|
+
{
|
|
179
|
+
"--disaggregation-transfer-backend": "nixl",
|
|
180
|
+
"--host": "::",
|
|
181
|
+
},
|
|
182
|
+
)
|
|
183
|
+
engine.extra_server_args = {
|
|
184
|
+
key: value
|
|
185
|
+
for key, value in engine.extra_server_args.items()
|
|
186
|
+
if _flag_name(key) != "--disaggregation-mode"
|
|
187
|
+
}
|
|
188
|
+
engine.base_url = _local_base_url("::", engine.worker_port)
|
|
189
|
+
self.router = (
|
|
190
|
+
_PDRouter(
|
|
191
|
+
pd_config=list(zip(roles, hosts)),
|
|
192
|
+
worker_port=engine.worker_port,
|
|
193
|
+
router_port=router_port,
|
|
194
|
+
prefill_bootstrap_port=engine.prefill_bootstrap_port,
|
|
195
|
+
health_timeout=engine.health_timeout,
|
|
196
|
+
api_key=api_key,
|
|
197
|
+
max_concurrent_requests=max_concurrent_requests,
|
|
198
|
+
drain_timeout=drain_timeout,
|
|
199
|
+
health_interval=health_interval,
|
|
200
|
+
health_request_timeout=router_health_request_timeout,
|
|
201
|
+
health_failure_threshold=health_failure_threshold,
|
|
202
|
+
)
|
|
203
|
+
if cluster.rank == 0
|
|
204
|
+
else None
|
|
205
|
+
)
|
|
206
|
+
self.drain_timeout = drain_timeout
|
|
207
|
+
self.health_interval = health_interval
|
|
208
|
+
self.health_failure_timeout = health_failure_timeout
|
|
209
|
+
self._stopped = threading.Event()
|
|
210
|
+
self._lock = threading.Lock()
|
|
211
|
+
self._started = False
|
|
212
|
+
self._threads: list[threading.Thread] = []
|
|
213
|
+
|
|
214
|
+
def start(self, *, warmup: Callable[[], None] | None = None) -> None:
|
|
215
|
+
"""Run warmup on rank zero before supervising every worker's health."""
|
|
216
|
+
if self._started or self._stopped.is_set():
|
|
217
|
+
raise RuntimeError("PDEndpoint can only be started once")
|
|
218
|
+
self._started = True
|
|
219
|
+
os.environ.setdefault("SGLANG_HOST_IP", self._host_ip)
|
|
220
|
+
_configure_fabric()
|
|
221
|
+
os.environ.update(
|
|
222
|
+
SGLANG_GRACEFUL_SHUTDOWN_TIMEOUT=str(self.drain_timeout),
|
|
223
|
+
SGLANG_ENABLE_HEALTH_ENDPOINT_GENERATION="1",
|
|
224
|
+
)
|
|
225
|
+
try:
|
|
226
|
+
for endpoint in (self.engine, self.router):
|
|
227
|
+
if endpoint is not None:
|
|
228
|
+
endpoint.start()
|
|
229
|
+
self._watch_process(endpoint)
|
|
230
|
+
if self.router is not None:
|
|
231
|
+
if warmup is not None:
|
|
232
|
+
warmup()
|
|
233
|
+
self.router.wait_ready()
|
|
234
|
+
for _, host in self.router.pd_config:
|
|
235
|
+
self._watch_health(_url(host, self.engine.worker_port) + "/health")
|
|
236
|
+
self._watch_health(self.router.base_url + "/health")
|
|
237
|
+
except BaseException:
|
|
238
|
+
self.stop()
|
|
239
|
+
raise
|
|
240
|
+
|
|
241
|
+
def _watch_process(self, endpoint: SGLangEndpoint | RouterEndpoint) -> None:
|
|
242
|
+
process = endpoint._proc
|
|
243
|
+
assert process is not None
|
|
244
|
+
|
|
245
|
+
def wait():
|
|
246
|
+
code = process.wait()
|
|
247
|
+
if not self._stopped.is_set():
|
|
248
|
+
print(f"[pd] Serving process exited with code {code}", flush=True)
|
|
249
|
+
self._fail()
|
|
250
|
+
|
|
251
|
+
thread = threading.Thread(target=wait, daemon=True)
|
|
252
|
+
self._threads.append(thread)
|
|
253
|
+
thread.start()
|
|
254
|
+
|
|
255
|
+
def _watch_health(self, url: str) -> None:
|
|
256
|
+
def poll():
|
|
257
|
+
failed_since = None
|
|
258
|
+
while not self._stopped.wait(self.health_interval):
|
|
259
|
+
started = time.monotonic()
|
|
260
|
+
error = _health_check(
|
|
261
|
+
url, request_timeout=self.engine.health_request_timeout
|
|
262
|
+
)
|
|
263
|
+
if error is None:
|
|
264
|
+
failed_since = None
|
|
265
|
+
continue
|
|
266
|
+
if failed_since is None:
|
|
267
|
+
failed_since = started
|
|
268
|
+
elapsed = time.monotonic() - failed_since
|
|
269
|
+
print(f"[pd] {url}: {error}; unhealthy for {elapsed:.1f}s", flush=True)
|
|
270
|
+
if elapsed >= self.health_failure_timeout:
|
|
271
|
+
self._fail()
|
|
272
|
+
return
|
|
273
|
+
|
|
274
|
+
thread = threading.Thread(target=poll, daemon=True)
|
|
275
|
+
self._threads.append(thread)
|
|
276
|
+
thread.start()
|
|
277
|
+
|
|
278
|
+
def _fail(self) -> None:
|
|
279
|
+
with self._lock:
|
|
280
|
+
if self._stopped.is_set():
|
|
281
|
+
return
|
|
282
|
+
self._stopped.set()
|
|
283
|
+
try:
|
|
284
|
+
for endpoint in (self.router, self.engine):
|
|
285
|
+
if endpoint is not None and endpoint._proc is not None:
|
|
286
|
+
try:
|
|
287
|
+
endpoint._proc.kill()
|
|
288
|
+
except ProcessLookupError:
|
|
289
|
+
pass
|
|
290
|
+
finally:
|
|
291
|
+
_terminate_container()
|
|
292
|
+
|
|
293
|
+
def stop(self) -> None:
|
|
294
|
+
with self._lock:
|
|
295
|
+
if self._stopped.is_set():
|
|
296
|
+
return
|
|
297
|
+
self._stopped.set()
|
|
298
|
+
deadline = time.monotonic() + self.drain_timeout + 30
|
|
299
|
+
try:
|
|
300
|
+
for endpoint in (self.router, self.engine):
|
|
301
|
+
if endpoint is None:
|
|
302
|
+
continue
|
|
303
|
+
process = endpoint._proc
|
|
304
|
+
if process is not None and process.poll() is None:
|
|
305
|
+
process.terminate()
|
|
306
|
+
try:
|
|
307
|
+
process.wait(timeout=max(0, deadline - time.monotonic()))
|
|
308
|
+
except subprocess.TimeoutExpired:
|
|
309
|
+
process.kill()
|
|
310
|
+
endpoint.stop()
|
|
311
|
+
finally:
|
|
312
|
+
for thread in self._threads:
|
|
313
|
+
thread.join(
|
|
314
|
+
timeout=self.engine.health_request_timeout + self.health_interval
|
|
315
|
+
)
|
|
@@ -1,15 +0,0 @@
|
|
|
1
|
-
# autoinference-utils
|
|
2
|
-
|
|
3
|
-
Shared endpoint abstractions (`SGLangEndpoint`, `VLLMEndpoint`) for autoinference deployments.
|
|
4
|
-
|
|
5
|
-
For SGLang models whose custom Transformers code cannot load from symlinked
|
|
6
|
-
Hugging Face snapshots, set `AUTOINFERENCE_MATERIALIZE_MODEL_PATH=1`. The endpoint
|
|
7
|
-
copies local base/draft model code and metadata into unique temporary directories
|
|
8
|
-
and symlinks top-level safetensors weights. Repository IDs pass through unchanged.
|
|
9
|
-
Leave the setting unset to preserve the original paths.
|
|
10
|
-
|
|
11
|
-
## Publishing
|
|
12
|
-
|
|
13
|
-
Bump `version` in `pyproject.toml` and merge to `main` — CI publishes automatically via trusted publishing.
|
|
14
|
-
|
|
15
|
-
To publish manually: `cd autoinference_utils && uv build && uv publish`
|
|
@@ -1,38 +0,0 @@
|
|
|
1
|
-
"""Prepare local model snapshots for Transformers' remote-code loader."""
|
|
2
|
-
|
|
3
|
-
import os
|
|
4
|
-
import shutil
|
|
5
|
-
import tempfile
|
|
6
|
-
from pathlib import Path
|
|
7
|
-
|
|
8
|
-
MATERIALIZE_MODEL_PATH_ENV = "AUTOINFERENCE_MATERIALIZE_MODEL_PATH"
|
|
9
|
-
|
|
10
|
-
|
|
11
|
-
def materialize_model_path(path: str) -> str:
|
|
12
|
-
"""Copy snapshot metadata/code into a real directory; retain weight links.
|
|
13
|
-
|
|
14
|
-
Transformers can resolve Python-module symlinks into the HF blob directory,
|
|
15
|
-
losing the relative imports next to the original snapshot entry. This opt-in
|
|
16
|
-
keeps small files together without copying large safetensors shards. The
|
|
17
|
-
unique directory lasts for the container's lifetime.
|
|
18
|
-
"""
|
|
19
|
-
source = Path(path)
|
|
20
|
-
if os.environ.get(MATERIALIZE_MODEL_PATH_ENV) != "1" or not source.is_dir():
|
|
21
|
-
return path
|
|
22
|
-
|
|
23
|
-
source = source.resolve()
|
|
24
|
-
destination = Path(tempfile.mkdtemp(prefix="autoinference-model-"))
|
|
25
|
-
try:
|
|
26
|
-
for entry in source.iterdir():
|
|
27
|
-
target = destination / entry.name
|
|
28
|
-
if entry.is_file() and entry.suffix == ".safetensors":
|
|
29
|
-
target.symlink_to(entry.resolve(strict=True))
|
|
30
|
-
elif entry.is_dir():
|
|
31
|
-
shutil.copytree(entry, target, symlinks=False)
|
|
32
|
-
else:
|
|
33
|
-
shutil.copy2(entry, target, follow_symlinks=True)
|
|
34
|
-
except BaseException:
|
|
35
|
-
shutil.rmtree(destination)
|
|
36
|
-
raise
|
|
37
|
-
print(f"Materialized model snapshot {source} -> {destination}")
|
|
38
|
-
return str(destination)
|
|
@@ -1,84 +0,0 @@
|
|
|
1
|
-
import shutil
|
|
2
|
-
from pathlib import Path
|
|
3
|
-
|
|
4
|
-
import pytest
|
|
5
|
-
|
|
6
|
-
from autoinference_utils.endpoint import SGLangEndpoint
|
|
7
|
-
from autoinference_utils.model_path import (
|
|
8
|
-
MATERIALIZE_MODEL_PATH_ENV,
|
|
9
|
-
materialize_model_path,
|
|
10
|
-
)
|
|
11
|
-
|
|
12
|
-
|
|
13
|
-
def snapshot(tmp_path):
|
|
14
|
-
blobs = tmp_path / "blobs"
|
|
15
|
-
source = tmp_path / "snapshots" / "revision"
|
|
16
|
-
blobs.mkdir()
|
|
17
|
-
source.mkdir(parents=True)
|
|
18
|
-
for name, content in {
|
|
19
|
-
"config.json": '{"model_type": "kimi"}',
|
|
20
|
-
"configuration_kimi.py": "from .helper import VALUE\n",
|
|
21
|
-
"helper.py": "VALUE = 42\n",
|
|
22
|
-
"model.safetensors": "weight placeholder",
|
|
23
|
-
}.items():
|
|
24
|
-
blob = blobs / ("hash-" + name)
|
|
25
|
-
blob.write_text(content)
|
|
26
|
-
(source / name).symlink_to(blob)
|
|
27
|
-
return source
|
|
28
|
-
|
|
29
|
-
|
|
30
|
-
def test_opt_in_and_repository_ids(tmp_path, monkeypatch):
|
|
31
|
-
source = snapshot(tmp_path)
|
|
32
|
-
monkeypatch.delenv(MATERIALIZE_MODEL_PATH_ENV, raising=False)
|
|
33
|
-
assert materialize_model_path(str(source)) == str(source)
|
|
34
|
-
monkeypatch.setenv(MATERIALIZE_MODEL_PATH_ENV, "1")
|
|
35
|
-
assert materialize_model_path("moonshotai/Kimi-K3") == "moonshotai/Kimi-K3"
|
|
36
|
-
|
|
37
|
-
|
|
38
|
-
def test_local_code_is_materialized_and_weights_stay_linked(tmp_path, monkeypatch):
|
|
39
|
-
source = snapshot(tmp_path)
|
|
40
|
-
monkeypatch.setenv(MATERIALIZE_MODEL_PATH_ENV, "1")
|
|
41
|
-
monkeypatch.setattr("tempfile.tempdir", str(tmp_path))
|
|
42
|
-
result = Path(materialize_model_path(str(source)))
|
|
43
|
-
assert result != source
|
|
44
|
-
for name in ("config.json", "configuration_kimi.py", "helper.py"):
|
|
45
|
-
assert not (result / name).is_symlink()
|
|
46
|
-
assert (result / name).resolve().parent == result
|
|
47
|
-
assert (result / name).read_bytes() == (source / name).read_bytes()
|
|
48
|
-
assert (result / "model.safetensors").is_symlink()
|
|
49
|
-
assert (result / "model.safetensors").resolve() == (
|
|
50
|
-
source / "model.safetensors"
|
|
51
|
-
).resolve()
|
|
52
|
-
assert (source / "configuration_kimi.py").is_symlink()
|
|
53
|
-
|
|
54
|
-
|
|
55
|
-
def test_endpoint_uses_materialized_base_and_draft(tmp_path, monkeypatch):
|
|
56
|
-
source = snapshot(tmp_path)
|
|
57
|
-
monkeypatch.setenv(MATERIALIZE_MODEL_PATH_ENV, "1")
|
|
58
|
-
monkeypatch.setattr("tempfile.tempdir", str(tmp_path))
|
|
59
|
-
endpoint = SGLangEndpoint(
|
|
60
|
-
model_path=str(source), speculative_model_path=str(source)
|
|
61
|
-
)
|
|
62
|
-
assert endpoint.model_path != str(source)
|
|
63
|
-
assert endpoint.speculative_model_path != str(source)
|
|
64
|
-
assert endpoint.model_path != endpoint.speculative_model_path
|
|
65
|
-
cmd = endpoint._build_cmd()
|
|
66
|
-
assert cmd[cmd.index("--model-path") + 1] == endpoint.model_path
|
|
67
|
-
assert (
|
|
68
|
-
cmd[cmd.index("--speculative-draft-model-path") + 1]
|
|
69
|
-
== endpoint.speculative_model_path
|
|
70
|
-
)
|
|
71
|
-
|
|
72
|
-
|
|
73
|
-
def test_failed_copy_removes_partial_directory(tmp_path, monkeypatch):
|
|
74
|
-
source = snapshot(tmp_path)
|
|
75
|
-
monkeypatch.setenv(MATERIALIZE_MODEL_PATH_ENV, "1")
|
|
76
|
-
monkeypatch.setattr("tempfile.tempdir", str(tmp_path))
|
|
77
|
-
|
|
78
|
-
def fail(*args, **kwargs):
|
|
79
|
-
raise OSError("copy failed")
|
|
80
|
-
|
|
81
|
-
monkeypatch.setattr(shutil, "copy2", fail)
|
|
82
|
-
with pytest.raises(OSError, match="copy failed"):
|
|
83
|
-
materialize_model_path(str(source))
|
|
84
|
-
assert not list(tmp_path.glob("autoinference-model-*"))
|
|
File without changes
|
|
File without changes
|