autoinference-utils 0.2.7__tar.gz → 0.2.9__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.9}/PKG-INFO +1 -7
- autoinference_utils-0.2.9/README.md +9 -0
- {autoinference_utils-0.2.7 → autoinference_utils-0.2.9}/pyproject.toml +1 -1
- {autoinference_utils-0.2.7 → autoinference_utils-0.2.9}/src/autoinference_utils/endpoint.py +51 -24
- autoinference_utils-0.2.9/src/autoinference_utils/pd.py +372 -0
- autoinference_utils-0.2.9/src/autoinference_utils/router.py +103 -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.9}/.gitignore +0 -0
- {autoinference_utils-0.2.7 → autoinference_utils-0.2.9}/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.9
|
|
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,8 +23,6 @@ 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
|
|
@@ -35,6 +33,14 @@ MODAL_SESSION_ID_HEADER = "modal-session-id"
|
|
|
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
|
|
|
@@ -150,18 +156,16 @@ class SGLangEndpoint(Endpoint):
|
|
|
150
156
|
sglang_port = (
|
|
151
157
|
worker_port + BENCH_MODE_PORT_OFFSET if self.bench_mode else worker_port
|
|
152
158
|
)
|
|
153
|
-
|
|
159
|
+
host_arg = get_server_arg(extra_server_args or {}, "--host")
|
|
160
|
+
host = (host_arg[1] or host_arg[0].partition("=")[2]) if host_arg else "0.0.0.0"
|
|
161
|
+
super().__init__(base_url=_local_base_url(host, sglang_port))
|
|
154
162
|
|
|
155
163
|
self.worker_port = sglang_port
|
|
156
|
-
self.model_path =
|
|
164
|
+
self.model_path = model_path
|
|
157
165
|
self.tp = tp
|
|
158
166
|
self.ep = ep
|
|
159
167
|
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
|
-
)
|
|
168
|
+
self.speculative_model_path = speculative_model_path
|
|
165
169
|
self.load_format = load_format
|
|
166
170
|
self.nnodes = nnodes
|
|
167
171
|
self.node_rank = node_rank
|
|
@@ -286,7 +290,7 @@ class SGLangEndpoint(Endpoint):
|
|
|
286
290
|
return cmd
|
|
287
291
|
|
|
288
292
|
def health_check(self) -> str | None:
|
|
289
|
-
url = f"
|
|
293
|
+
url = f"{self.base_url}/health"
|
|
290
294
|
return _health_check(
|
|
291
295
|
url,
|
|
292
296
|
request_timeout=self.health_request_timeout,
|
|
@@ -300,6 +304,7 @@ class SGLangEndpoint(Endpoint):
|
|
|
300
304
|
wait_ready(
|
|
301
305
|
self._proc,
|
|
302
306
|
port=self.worker_port,
|
|
307
|
+
base_url=self.base_url,
|
|
303
308
|
timeout=self.health_timeout,
|
|
304
309
|
poll_interval=self.health_poll_interval,
|
|
305
310
|
request_timeout=self.health_request_timeout,
|
|
@@ -308,6 +313,7 @@ class SGLangEndpoint(Endpoint):
|
|
|
308
313
|
self._bench_server = start_bench_proxy(
|
|
309
314
|
listen_port=self.listen_port,
|
|
310
315
|
upstream_port=self.worker_port,
|
|
316
|
+
upstream_base_url=self.base_url,
|
|
311
317
|
)
|
|
312
318
|
|
|
313
319
|
def stop(self):
|
|
@@ -474,13 +480,16 @@ class RouterEndpoint(Endpoint):
|
|
|
474
480
|
api_key: Optional[str] = None,
|
|
475
481
|
health_timeout: float = 10 * 60,
|
|
476
482
|
health_poll_interval: float = 5.0,
|
|
483
|
+
health_path: str = "/health",
|
|
484
|
+
host: str = "0.0.0.0",
|
|
477
485
|
):
|
|
478
486
|
self.bench_mode = os.environ.get(BENCH_MODE_ENV) == "1"
|
|
479
487
|
self.listen_port = router_port
|
|
480
488
|
actual_router_port = (
|
|
481
489
|
router_port + BENCH_MODE_PORT_OFFSET if self.bench_mode else router_port
|
|
482
490
|
)
|
|
483
|
-
|
|
491
|
+
self.host = host
|
|
492
|
+
super().__init__(base_url=_local_base_url(host, actual_router_port))
|
|
484
493
|
self.pd_config = list(pd_config)
|
|
485
494
|
self.worker_port = (
|
|
486
495
|
worker_port + BENCH_MODE_PORT_OFFSET if self.bench_mode else worker_port
|
|
@@ -490,6 +499,7 @@ class RouterEndpoint(Endpoint):
|
|
|
490
499
|
self.api_key = api_key
|
|
491
500
|
self.health_timeout = health_timeout
|
|
492
501
|
self.health_poll_interval = health_poll_interval
|
|
502
|
+
self.health_path = health_path
|
|
493
503
|
self._proc: Optional[subprocess.Popen] = None
|
|
494
504
|
self._bench_server: Optional[http.server.ThreadingHTTPServer] = None
|
|
495
505
|
|
|
@@ -499,7 +509,7 @@ class RouterEndpoint(Endpoint):
|
|
|
499
509
|
"-m",
|
|
500
510
|
"sglang_router.launch_router",
|
|
501
511
|
"--host",
|
|
502
|
-
"
|
|
512
|
+
f"[{self.host}]" if ":" in self.host else self.host,
|
|
503
513
|
"--port",
|
|
504
514
|
str(self.router_port),
|
|
505
515
|
"--prefill-policy",
|
|
@@ -528,7 +538,7 @@ class RouterEndpoint(Endpoint):
|
|
|
528
538
|
cmd.append("--pd-disaggregation")
|
|
529
539
|
|
|
530
540
|
for role, node_ip in self.pd_config:
|
|
531
|
-
node_url =
|
|
541
|
+
node_url = _url(node_ip, self.worker_port)
|
|
532
542
|
if role == "prefill":
|
|
533
543
|
cmd.extend(["--prefill", node_url, str(self.prefill_bootstrap_port)])
|
|
534
544
|
elif role == "decode":
|
|
@@ -543,7 +553,7 @@ class RouterEndpoint(Endpoint):
|
|
|
543
553
|
def start(self):
|
|
544
554
|
for _, node_ip in self.pd_config:
|
|
545
555
|
_wait_ready_url(
|
|
546
|
-
f"
|
|
556
|
+
f"{_url(node_ip, self.worker_port)}/health",
|
|
547
557
|
timeout=self.health_timeout,
|
|
548
558
|
poll_interval=self.health_poll_interval,
|
|
549
559
|
)
|
|
@@ -551,17 +561,25 @@ class RouterEndpoint(Endpoint):
|
|
|
551
561
|
cmd = self._build_cmd()
|
|
552
562
|
print(f"[router] starting: {shlex.join(cmd)}")
|
|
553
563
|
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
|
-
)
|
|
564
|
+
self.wait_ready()
|
|
559
565
|
if self.bench_mode:
|
|
560
566
|
self._bench_server = start_bench_proxy(
|
|
561
567
|
listen_port=self.listen_port,
|
|
562
568
|
upstream_port=self.router_port,
|
|
569
|
+
upstream_base_url=self.base_url,
|
|
563
570
|
)
|
|
564
571
|
|
|
572
|
+
def wait_ready(self) -> None:
|
|
573
|
+
assert self._proc is not None
|
|
574
|
+
wait_ready(
|
|
575
|
+
self._proc,
|
|
576
|
+
port=self.router_port,
|
|
577
|
+
base_url=self.base_url,
|
|
578
|
+
health_path=self.health_path,
|
|
579
|
+
timeout=self.health_timeout,
|
|
580
|
+
poll_interval=self.health_poll_interval,
|
|
581
|
+
)
|
|
582
|
+
|
|
565
583
|
def stop(self):
|
|
566
584
|
if self._bench_server is not None:
|
|
567
585
|
self._bench_server.shutdown()
|
|
@@ -579,11 +597,12 @@ def start_bench_proxy(
|
|
|
579
597
|
*,
|
|
580
598
|
listen_port: int,
|
|
581
599
|
upstream_port: int,
|
|
600
|
+
upstream_base_url: str | None = None,
|
|
582
601
|
) -> http.server.ThreadingHTTPServer:
|
|
583
602
|
handler_cls = type(
|
|
584
603
|
"BenchProxyHandler",
|
|
585
604
|
(_BenchProxyHandler,),
|
|
586
|
-
{"upstream_port": upstream_port},
|
|
605
|
+
{"upstream_port": upstream_port, "upstream_base_url": upstream_base_url},
|
|
587
606
|
)
|
|
588
607
|
server = http.server.ThreadingHTTPServer(("0.0.0.0", listen_port), handler_cls)
|
|
589
608
|
thread = threading.Thread(
|
|
@@ -596,6 +615,7 @@ def start_bench_proxy(
|
|
|
596
615
|
|
|
597
616
|
class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
|
|
598
617
|
upstream_port: int = 8000
|
|
618
|
+
upstream_base_url: str | None = None
|
|
599
619
|
|
|
600
620
|
def log_message(self, format, *args):
|
|
601
621
|
return
|
|
@@ -623,7 +643,11 @@ class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
|
|
|
623
643
|
|
|
624
644
|
benchmark = payload.get("benchmark") or ""
|
|
625
645
|
args = payload.get("args") or []
|
|
626
|
-
target =
|
|
646
|
+
target = (
|
|
647
|
+
payload.get("target")
|
|
648
|
+
or self.upstream_base_url
|
|
649
|
+
or f"http://localhost:{self.upstream_port}"
|
|
650
|
+
)
|
|
627
651
|
output_dir = payload.get("output_dir") or "/tmp/bench-output"
|
|
628
652
|
|
|
629
653
|
if not benchmark:
|
|
@@ -636,7 +660,8 @@ class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
|
|
|
636
660
|
self._send_raw(status, "application/json", body)
|
|
637
661
|
|
|
638
662
|
def _proxy(self, method: str):
|
|
639
|
-
|
|
663
|
+
base_url = self.upstream_base_url or f"http://localhost:{self.upstream_port}"
|
|
664
|
+
url = f"{base_url}{self.path}"
|
|
640
665
|
length = int(self.headers.get("Content-Length", 0))
|
|
641
666
|
body = self.rfile.read(length) if length else None
|
|
642
667
|
forward_headers = {
|
|
@@ -759,12 +784,13 @@ def wait_ready(
|
|
|
759
784
|
port: int,
|
|
760
785
|
timeout: float,
|
|
761
786
|
health_path: str = "/health",
|
|
787
|
+
base_url: str | None = None,
|
|
762
788
|
poll_interval: float = 5.0,
|
|
763
789
|
request_timeout: float = 5.0,
|
|
764
790
|
) -> None:
|
|
765
791
|
"""Poll an HTTP health endpoint until ready, raising if the process dies."""
|
|
766
792
|
deadline = time.time() + timeout
|
|
767
|
-
url = f"
|
|
793
|
+
url = f"{base_url or _local_base_url('127.0.0.1', port)}{health_path}"
|
|
768
794
|
last_error = "no response yet"
|
|
769
795
|
|
|
770
796
|
while time.time() < deadline:
|
|
@@ -794,6 +820,7 @@ def wait_ready(
|
|
|
794
820
|
def warmup_chat_completions(
|
|
795
821
|
*,
|
|
796
822
|
port: int,
|
|
823
|
+
base_url: str | None = None,
|
|
797
824
|
payload: Mapping[str, Any],
|
|
798
825
|
headers: Mapping[str, str] | None = None,
|
|
799
826
|
successful_requests: int = 3,
|
|
@@ -802,7 +829,7 @@ def warmup_chat_completions(
|
|
|
802
829
|
retry_delay: float = 1.0,
|
|
803
830
|
) -> None:
|
|
804
831
|
"""Warm the OpenAI chat completions endpoint with strict retries."""
|
|
805
|
-
url = f"
|
|
832
|
+
url = f"{base_url or _local_base_url('127.0.0.1', port)}/v1/chat/completions"
|
|
806
833
|
request_headers = {"Content-Type": "application/json"}
|
|
807
834
|
if headers:
|
|
808
835
|
request_headers.update(headers)
|
|
@@ -0,0 +1,372 @@
|
|
|
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
|
+
server_arg_tokens,
|
|
23
|
+
)
|
|
24
|
+
from .router import verify_router
|
|
25
|
+
|
|
26
|
+
# PDEndpoint owns the topology and every flag it mirrors in its own supervision
|
|
27
|
+
# (drain deadline, health probes); recipes tune everything else via router_args.
|
|
28
|
+
_MANAGED_ROUTER_FLAGS = frozenset(
|
|
29
|
+
{
|
|
30
|
+
"--host",
|
|
31
|
+
"--port",
|
|
32
|
+
"--pd-disaggregation",
|
|
33
|
+
"--prefill",
|
|
34
|
+
"--decode",
|
|
35
|
+
"--api-key",
|
|
36
|
+
"--shutdown-grace-period-secs",
|
|
37
|
+
"--health-check-interval-secs",
|
|
38
|
+
"--health-check-timeout-secs",
|
|
39
|
+
"--health-failure-threshold",
|
|
40
|
+
}
|
|
41
|
+
)
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def _managed_router_tokens(router_args: Mapping[str, str]) -> list[str]:
|
|
45
|
+
"""Rendered override tokens that would reach a managed flag.
|
|
46
|
+
|
|
47
|
+
Checks argv as the router parses it: values split into extra tokens, and
|
|
48
|
+
argparse accepts any unambiguous prefix of a long flag.
|
|
49
|
+
"""
|
|
50
|
+
tokens = [
|
|
51
|
+
token
|
|
52
|
+
for key, value in router_args.items()
|
|
53
|
+
for token in server_arg_tokens(key, value)
|
|
54
|
+
]
|
|
55
|
+
return sorted(
|
|
56
|
+
{
|
|
57
|
+
token
|
|
58
|
+
for token in tokens
|
|
59
|
+
if token.startswith("-")
|
|
60
|
+
and any(
|
|
61
|
+
flag.startswith(_flag_name(token)) for flag in _MANAGED_ROUTER_FLAGS
|
|
62
|
+
)
|
|
63
|
+
}
|
|
64
|
+
)
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
def _configure_fabric() -> None:
|
|
68
|
+
if os.environ.get("FI_PROVIDER") == "efa":
|
|
69
|
+
os.environ["SGLANG_DISAGGREGATION_NIXL_BACKEND"] = "LIBFABRIC"
|
|
70
|
+
elif hca := os.environ.get("NCCL_IB_HCA", ""):
|
|
71
|
+
rails = [name for name in hca.lstrip("=").split(",") if name]
|
|
72
|
+
os.environ["UCX_NET_DEVICES"] = ",".join(f"{name}:1" for name in rails)
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
async def _stop_container() -> None:
|
|
76
|
+
from modal.client import _Client
|
|
77
|
+
from modal_proto import api_pb2
|
|
78
|
+
|
|
79
|
+
async def stop():
|
|
80
|
+
client = await _Client.from_env()
|
|
81
|
+
await client.stub.ContainerStop(
|
|
82
|
+
api_pb2.ContainerStopRequest(task_id=os.environ["MODAL_TASK_ID"])
|
|
83
|
+
)
|
|
84
|
+
|
|
85
|
+
await asyncio.wait_for(stop(), timeout=3)
|
|
86
|
+
|
|
87
|
+
|
|
88
|
+
def _terminate_container() -> None:
|
|
89
|
+
from modal._utils.async_utils import synchronize_api
|
|
90
|
+
|
|
91
|
+
try:
|
|
92
|
+
synchronize_api(_stop_container)()
|
|
93
|
+
finally:
|
|
94
|
+
os._exit(1)
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
class _PDRouter(RouterEndpoint):
|
|
98
|
+
def __init__(
|
|
99
|
+
self,
|
|
100
|
+
*,
|
|
101
|
+
max_concurrent_requests,
|
|
102
|
+
drain_timeout,
|
|
103
|
+
health_interval,
|
|
104
|
+
health_request_timeout,
|
|
105
|
+
health_failure_threshold,
|
|
106
|
+
router_args,
|
|
107
|
+
**kwargs,
|
|
108
|
+
):
|
|
109
|
+
super().__init__(health_path="/readiness", host="::", **kwargs)
|
|
110
|
+
self.max_concurrent_requests = max_concurrent_requests
|
|
111
|
+
self.drain_timeout = drain_timeout
|
|
112
|
+
self.health_interval = health_interval
|
|
113
|
+
self.health_request_timeout = health_request_timeout
|
|
114
|
+
self.health_failure_threshold = health_failure_threshold
|
|
115
|
+
self.router_args = dict(router_args)
|
|
116
|
+
|
|
117
|
+
def _build_cmd(self) -> list[str]:
|
|
118
|
+
cmd = [
|
|
119
|
+
"python",
|
|
120
|
+
"-m",
|
|
121
|
+
"sglang_router.launch_router",
|
|
122
|
+
"--host",
|
|
123
|
+
"[::]",
|
|
124
|
+
"--port",
|
|
125
|
+
str(self.router_port),
|
|
126
|
+
"--pd-disaggregation",
|
|
127
|
+
]
|
|
128
|
+
for role, host in self.pd_config:
|
|
129
|
+
cmd.extend([f"--{role}", _url(host, self.worker_port)])
|
|
130
|
+
if role == "prefill":
|
|
131
|
+
cmd.append(str(self.prefill_bootstrap_port))
|
|
132
|
+
defaults = {
|
|
133
|
+
"--prefill-policy": "cache_aware",
|
|
134
|
+
"--decode-policy": "round_robin",
|
|
135
|
+
"--max-concurrent-requests": str(self.max_concurrent_requests),
|
|
136
|
+
"--queue-size": "0",
|
|
137
|
+
"--rate-limit-tokens-per-second": "0",
|
|
138
|
+
"--health-check-interval-secs": str(self.health_interval),
|
|
139
|
+
"--health-check-timeout-secs": str(self.health_request_timeout),
|
|
140
|
+
"--health-failure-threshold": str(self.health_failure_threshold),
|
|
141
|
+
"--shutdown-grace-period-secs": str(self.drain_timeout),
|
|
142
|
+
"--request-timeout-secs": "1800",
|
|
143
|
+
"--log-level": "info",
|
|
144
|
+
}
|
|
145
|
+
for flag, value in _merge_server_args(defaults, self.router_args).items():
|
|
146
|
+
cmd.extend(server_arg_tokens(flag, value))
|
|
147
|
+
if self.api_key is not None:
|
|
148
|
+
cmd.extend(["--api-key", self.api_key])
|
|
149
|
+
return cmd
|
|
150
|
+
|
|
151
|
+
|
|
152
|
+
class PDEndpoint(Endpoint):
|
|
153
|
+
"""Own a local SGLang engine and the rank-zero router for P/D serving.
|
|
154
|
+
|
|
155
|
+
Call start/stop from Modal enter/exit hooks. Model-specific arguments and
|
|
156
|
+
warmup stay in the recipe. The image must support NIXL and graceful shutdown.
|
|
157
|
+
The ratio is (prefill, decode) counts; prefill ranks precede decode ranks.
|
|
158
|
+
The router is the build pinned in autoinference_utils.router; router_args
|
|
159
|
+
override its default flags, and allow_custom_router skips the pin check.
|
|
160
|
+
"""
|
|
161
|
+
|
|
162
|
+
def __init__(
|
|
163
|
+
self,
|
|
164
|
+
engine: SGLangEndpoint,
|
|
165
|
+
*,
|
|
166
|
+
ratio: tuple[int, int],
|
|
167
|
+
prefill_args: Mapping[str, str] | None = None,
|
|
168
|
+
decode_args: Mapping[str, str] | None = None,
|
|
169
|
+
router_port: int = 9000,
|
|
170
|
+
max_concurrent_requests: int = 96,
|
|
171
|
+
api_key: str | None = None,
|
|
172
|
+
drain_timeout: int = 120,
|
|
173
|
+
health_interval: int = 10,
|
|
174
|
+
health_request_timeout: int = 10,
|
|
175
|
+
router_health_request_timeout: int = 5,
|
|
176
|
+
health_failure_threshold: int = 3,
|
|
177
|
+
health_failure_timeout: float = 30,
|
|
178
|
+
router_args: Mapping[str, str] | None = None,
|
|
179
|
+
allow_custom_router: bool = False,
|
|
180
|
+
):
|
|
181
|
+
from modal.experimental import get_cluster_info
|
|
182
|
+
|
|
183
|
+
if not isinstance(engine, SGLangEndpoint):
|
|
184
|
+
raise TypeError("PDEndpoint requires a SGLangEndpoint")
|
|
185
|
+
if engine._proc is not None or engine.bench_mode:
|
|
186
|
+
raise ValueError("PDEndpoint requires an unstarted, non-benchmark engine")
|
|
187
|
+
if (
|
|
188
|
+
min(
|
|
189
|
+
health_interval,
|
|
190
|
+
health_request_timeout,
|
|
191
|
+
router_health_request_timeout,
|
|
192
|
+
health_failure_threshold,
|
|
193
|
+
)
|
|
194
|
+
<= 0
|
|
195
|
+
or drain_timeout < 0
|
|
196
|
+
or not math.isfinite(health_failure_timeout)
|
|
197
|
+
or health_failure_timeout <= 0
|
|
198
|
+
):
|
|
199
|
+
raise ValueError(
|
|
200
|
+
"health settings must be positive and drain_timeout nonnegative"
|
|
201
|
+
)
|
|
202
|
+
if len(ratio) != 2 or any(
|
|
203
|
+
type(count) is not int or count < 1 for count in ratio
|
|
204
|
+
):
|
|
205
|
+
raise ValueError(
|
|
206
|
+
"ratio must contain positive integer prefill and decode counts"
|
|
207
|
+
)
|
|
208
|
+
if managed := _managed_router_tokens(router_args or {}):
|
|
209
|
+
raise ValueError(
|
|
210
|
+
f"router_args cannot set PDEndpoint-managed flags: {managed}"
|
|
211
|
+
)
|
|
212
|
+
roles = ("prefill",) * ratio[0] + ("decode",) * ratio[1]
|
|
213
|
+
cluster = get_cluster_info()
|
|
214
|
+
hosts = cluster.container_ips
|
|
215
|
+
if len(hosts) != len(roles) or not 0 <= cluster.rank < len(roles):
|
|
216
|
+
raise ValueError(
|
|
217
|
+
f"P/D ratio {ratio} requires a {len(roles)}-container Modal cluster"
|
|
218
|
+
)
|
|
219
|
+
self._host_ip = hosts[cluster.rank]
|
|
220
|
+
self.role = roles[cluster.rank]
|
|
221
|
+
super().__init__(_url(hosts[0], router_port))
|
|
222
|
+
self.engine = engine
|
|
223
|
+
engine.disaggregation_mode = self.role
|
|
224
|
+
engine.health_request_timeout = health_request_timeout
|
|
225
|
+
engine.extra_server_args = _merge_server_args(
|
|
226
|
+
engine.extra_server_args,
|
|
227
|
+
(prefill_args if self.role == "prefill" else decode_args) or {},
|
|
228
|
+
)
|
|
229
|
+
engine.extra_server_args = _merge_server_args(
|
|
230
|
+
engine.extra_server_args,
|
|
231
|
+
{
|
|
232
|
+
"--disaggregation-transfer-backend": "nixl",
|
|
233
|
+
"--host": "::",
|
|
234
|
+
},
|
|
235
|
+
)
|
|
236
|
+
engine.extra_server_args = {
|
|
237
|
+
key: value
|
|
238
|
+
for key, value in engine.extra_server_args.items()
|
|
239
|
+
if _flag_name(key) != "--disaggregation-mode"
|
|
240
|
+
}
|
|
241
|
+
engine.base_url = _local_base_url("::", engine.worker_port)
|
|
242
|
+
self.router = (
|
|
243
|
+
_PDRouter(
|
|
244
|
+
pd_config=list(zip(roles, hosts)),
|
|
245
|
+
worker_port=engine.worker_port,
|
|
246
|
+
router_port=router_port,
|
|
247
|
+
prefill_bootstrap_port=engine.prefill_bootstrap_port,
|
|
248
|
+
health_timeout=engine.health_timeout,
|
|
249
|
+
api_key=api_key,
|
|
250
|
+
max_concurrent_requests=max_concurrent_requests,
|
|
251
|
+
drain_timeout=drain_timeout,
|
|
252
|
+
health_interval=health_interval,
|
|
253
|
+
health_request_timeout=router_health_request_timeout,
|
|
254
|
+
health_failure_threshold=health_failure_threshold,
|
|
255
|
+
router_args=router_args or {},
|
|
256
|
+
)
|
|
257
|
+
if cluster.rank == 0
|
|
258
|
+
else None
|
|
259
|
+
)
|
|
260
|
+
# Fail before any weights load, not after a 20-minute engine boot.
|
|
261
|
+
if self.router is not None and not allow_custom_router:
|
|
262
|
+
verify_router()
|
|
263
|
+
self.drain_timeout = drain_timeout
|
|
264
|
+
self.health_interval = health_interval
|
|
265
|
+
self.health_failure_timeout = health_failure_timeout
|
|
266
|
+
self._stopped = threading.Event()
|
|
267
|
+
self._lock = threading.Lock()
|
|
268
|
+
self._started = False
|
|
269
|
+
self._threads: list[threading.Thread] = []
|
|
270
|
+
|
|
271
|
+
def start(self, *, warmup: Callable[[], None] | None = None) -> None:
|
|
272
|
+
"""Run warmup on rank zero before supervising every worker's health."""
|
|
273
|
+
if self._started or self._stopped.is_set():
|
|
274
|
+
raise RuntimeError("PDEndpoint can only be started once")
|
|
275
|
+
self._started = True
|
|
276
|
+
os.environ.setdefault("SGLANG_HOST_IP", self._host_ip)
|
|
277
|
+
_configure_fabric()
|
|
278
|
+
os.environ.update(
|
|
279
|
+
SGLANG_GRACEFUL_SHUTDOWN_TIMEOUT=str(self.drain_timeout),
|
|
280
|
+
SGLANG_ENABLE_HEALTH_ENDPOINT_GENERATION="1",
|
|
281
|
+
)
|
|
282
|
+
try:
|
|
283
|
+
for endpoint in (self.engine, self.router):
|
|
284
|
+
if endpoint is not None:
|
|
285
|
+
endpoint.start()
|
|
286
|
+
self._watch_process(endpoint)
|
|
287
|
+
if self.router is not None:
|
|
288
|
+
if warmup is not None:
|
|
289
|
+
warmup()
|
|
290
|
+
self.router.wait_ready()
|
|
291
|
+
for _, host in self.router.pd_config:
|
|
292
|
+
self._watch_health(_url(host, self.engine.worker_port) + "/health")
|
|
293
|
+
self._watch_health(self.router.base_url + "/health")
|
|
294
|
+
except BaseException:
|
|
295
|
+
self.stop()
|
|
296
|
+
raise
|
|
297
|
+
|
|
298
|
+
def _watch_process(self, endpoint: SGLangEndpoint | RouterEndpoint) -> None:
|
|
299
|
+
process = endpoint._proc
|
|
300
|
+
assert process is not None
|
|
301
|
+
|
|
302
|
+
def wait():
|
|
303
|
+
code = process.wait()
|
|
304
|
+
if not self._stopped.is_set():
|
|
305
|
+
print(f"[pd] Serving process exited with code {code}", flush=True)
|
|
306
|
+
self._fail()
|
|
307
|
+
|
|
308
|
+
thread = threading.Thread(target=wait, daemon=True)
|
|
309
|
+
self._threads.append(thread)
|
|
310
|
+
thread.start()
|
|
311
|
+
|
|
312
|
+
def _watch_health(self, url: str) -> None:
|
|
313
|
+
def poll():
|
|
314
|
+
failed_since = None
|
|
315
|
+
while not self._stopped.wait(self.health_interval):
|
|
316
|
+
started = time.monotonic()
|
|
317
|
+
error = _health_check(
|
|
318
|
+
url, request_timeout=self.engine.health_request_timeout
|
|
319
|
+
)
|
|
320
|
+
if error is None:
|
|
321
|
+
failed_since = None
|
|
322
|
+
continue
|
|
323
|
+
if failed_since is None:
|
|
324
|
+
failed_since = started
|
|
325
|
+
elapsed = time.monotonic() - failed_since
|
|
326
|
+
print(f"[pd] {url}: {error}; unhealthy for {elapsed:.1f}s", flush=True)
|
|
327
|
+
if elapsed >= self.health_failure_timeout:
|
|
328
|
+
self._fail()
|
|
329
|
+
return
|
|
330
|
+
|
|
331
|
+
thread = threading.Thread(target=poll, daemon=True)
|
|
332
|
+
self._threads.append(thread)
|
|
333
|
+
thread.start()
|
|
334
|
+
|
|
335
|
+
def _fail(self) -> None:
|
|
336
|
+
with self._lock:
|
|
337
|
+
if self._stopped.is_set():
|
|
338
|
+
return
|
|
339
|
+
self._stopped.set()
|
|
340
|
+
try:
|
|
341
|
+
for endpoint in (self.router, self.engine):
|
|
342
|
+
if endpoint is not None and endpoint._proc is not None:
|
|
343
|
+
try:
|
|
344
|
+
endpoint._proc.kill()
|
|
345
|
+
except ProcessLookupError:
|
|
346
|
+
pass
|
|
347
|
+
finally:
|
|
348
|
+
_terminate_container()
|
|
349
|
+
|
|
350
|
+
def stop(self) -> None:
|
|
351
|
+
with self._lock:
|
|
352
|
+
if self._stopped.is_set():
|
|
353
|
+
return
|
|
354
|
+
self._stopped.set()
|
|
355
|
+
deadline = time.monotonic() + self.drain_timeout + 30
|
|
356
|
+
try:
|
|
357
|
+
for endpoint in (self.router, self.engine):
|
|
358
|
+
if endpoint is None:
|
|
359
|
+
continue
|
|
360
|
+
process = endpoint._proc
|
|
361
|
+
if process is not None and process.poll() is None:
|
|
362
|
+
process.terminate()
|
|
363
|
+
try:
|
|
364
|
+
process.wait(timeout=max(0, deadline - time.monotonic()))
|
|
365
|
+
except subprocess.TimeoutExpired:
|
|
366
|
+
process.kill()
|
|
367
|
+
endpoint.stop()
|
|
368
|
+
finally:
|
|
369
|
+
for thread in self._threads:
|
|
370
|
+
thread.join(
|
|
371
|
+
timeout=self.engine.health_request_timeout + self.health_interval
|
|
372
|
+
)
|
|
@@ -0,0 +1,103 @@
|
|
|
1
|
+
"""The pinned SGLang model gateway (SMG) wheel that PDEndpoint routes through.
|
|
2
|
+
|
|
3
|
+
This module is the source of truth for the P/D router. The wheel is built once
|
|
4
|
+
from the fork commit below and published as a GitHub release asset; bumping
|
|
5
|
+
the pin here and releasing autoinference-utils moves every recipe that takes
|
|
6
|
+
the new release. Install it inside an image after autoinference-utils:
|
|
7
|
+
|
|
8
|
+
image.run_commands("python -m autoinference_utils.router install")
|
|
9
|
+
"""
|
|
10
|
+
|
|
11
|
+
from __future__ import annotations
|
|
12
|
+
|
|
13
|
+
import json
|
|
14
|
+
import subprocess
|
|
15
|
+
import sys
|
|
16
|
+
from importlib.metadata import Distribution, PackageNotFoundError, distribution
|
|
17
|
+
|
|
18
|
+
ROUTER_REPOSITORY = "https://github.com/modal-labs/sglang"
|
|
19
|
+
# Kimi K3 production router: HRRN decode preallocation (SMG_PD_DECODE_HRRN)
|
|
20
|
+
# and brace-free metrics HELP text (kimi-k3-sglang#158). The wheel is the one
|
|
21
|
+
# extracted from the Kimi K3 production image built at this commit.
|
|
22
|
+
ROUTER_COMMIT = "e3e602e3ae46fa4da7ec105d5dc3efcfc85c226b"
|
|
23
|
+
# The wheel is x86-64 Linux and links against glibc 2.39 (Ubuntu 24.04, which
|
|
24
|
+
# the lmsysorg/sglang and vllm base images use); uv refuses it on older images.
|
|
25
|
+
ROUTER_WHEEL_URL = (
|
|
26
|
+
f"{ROUTER_REPOSITORY}/releases/download/smg-router-{ROUTER_COMMIT[:8]}/"
|
|
27
|
+
"sglang_router-0.3.2-cp38-abi3-manylinux_2_39_x86_64.whl"
|
|
28
|
+
)
|
|
29
|
+
ROUTER_WHEEL_SHA256 = "5b867c6c308a609395a6f5c146023274f570b7187b3a7b92ebb11658da52a6c7"
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def router_requirement() -> str:
|
|
33
|
+
"""The hash-pinned direct requirement for the router wheel."""
|
|
34
|
+
return f"sglang-router @ {ROUTER_WHEEL_URL}#sha256={ROUTER_WHEEL_SHA256}"
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
def verify_router(dist: Distribution | None = None) -> None:
|
|
38
|
+
"""Raise unless the installed sglang-router came from the pinned wheel.
|
|
39
|
+
|
|
40
|
+
Checks the installer's PEP 610 record: the URL must match, and when the
|
|
41
|
+
installer recorded a digest it must match too. uv leaves archive_info empty
|
|
42
|
+
and enforces the digest at install time through the URL's sha256 fragment.
|
|
43
|
+
"""
|
|
44
|
+
try:
|
|
45
|
+
dist = dist or distribution("sglang-router")
|
|
46
|
+
origin = json.loads(dist.read_text("direct_url.json") or "{}")
|
|
47
|
+
except PackageNotFoundError:
|
|
48
|
+
origin = None
|
|
49
|
+
except ValueError:
|
|
50
|
+
origin = {}
|
|
51
|
+
if origin is None:
|
|
52
|
+
found = "no sglang-router is installed"
|
|
53
|
+
else:
|
|
54
|
+
recorded = ((origin.get("archive_info") or {}).get("hashes") or {}).get(
|
|
55
|
+
"sha256"
|
|
56
|
+
)
|
|
57
|
+
if origin.get("url") == ROUTER_WHEEL_URL and recorded in (
|
|
58
|
+
None,
|
|
59
|
+
ROUTER_WHEEL_SHA256,
|
|
60
|
+
):
|
|
61
|
+
return
|
|
62
|
+
found = f"installed sglang-router is {origin.get('url') or 'not a direct URL install'}"
|
|
63
|
+
if recorded is not None:
|
|
64
|
+
found += f" (sha256 {recorded})"
|
|
65
|
+
raise RuntimeError(
|
|
66
|
+
f"{found}; expected {ROUTER_WHEEL_URL} (sha256 {ROUTER_WHEEL_SHA256}). "
|
|
67
|
+
"Run `python -m autoinference_utils.router install` in the image after "
|
|
68
|
+
"autoinference-utils, or pass allow_custom_router=True."
|
|
69
|
+
)
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
def install() -> None:
|
|
73
|
+
"""Install the pinned wheel over whatever router the image ships."""
|
|
74
|
+
# --no-deps keeps the engine image's own web stack untouched.
|
|
75
|
+
subprocess.run(
|
|
76
|
+
[
|
|
77
|
+
"uv",
|
|
78
|
+
"pip",
|
|
79
|
+
"install",
|
|
80
|
+
"--python",
|
|
81
|
+
sys.executable,
|
|
82
|
+
"--no-cache",
|
|
83
|
+
"--no-deps",
|
|
84
|
+
"--reinstall-package",
|
|
85
|
+
"sglang-router",
|
|
86
|
+
router_requirement(),
|
|
87
|
+
],
|
|
88
|
+
check=True,
|
|
89
|
+
)
|
|
90
|
+
verify_router()
|
|
91
|
+
# --no-deps means a pin whose dependencies the image lacks must fail here,
|
|
92
|
+
# at image build, not at container start.
|
|
93
|
+
subprocess.run(
|
|
94
|
+
[sys.executable, "-m", "sglang_router.launch_router", "--help"],
|
|
95
|
+
check=True,
|
|
96
|
+
stdout=subprocess.DEVNULL,
|
|
97
|
+
)
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
if __name__ == "__main__":
|
|
101
|
+
if sys.argv[1:] != ["install"]:
|
|
102
|
+
raise SystemExit("usage: python -m autoinference_utils.router install")
|
|
103
|
+
install()
|
|
@@ -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
|