autoinference-utils 0.2.6__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.6 → autoinference_utils-0.2.8}/.gitignore +18 -1
- {autoinference_utils-0.2.6 → autoinference_utils-0.2.8}/PKG-INFO +1 -1
- {autoinference_utils-0.2.6 → autoinference_utils-0.2.8}/pyproject.toml +3 -3
- {autoinference_utils-0.2.6 → autoinference_utils-0.2.8/src}/autoinference_utils/endpoint.py +45 -15
- autoinference_utils-0.2.8/src/autoinference_utils/pd.py +315 -0
- {autoinference_utils-0.2.6 → autoinference_utils-0.2.8}/README.md +0 -0
- {autoinference_utils-0.2.6 → autoinference_utils-0.2.8/src}/autoinference_utils/__init__.py +0 -0
|
@@ -150,6 +150,7 @@ activemq-data/
|
|
|
150
150
|
|
|
151
151
|
# Environments
|
|
152
152
|
.env
|
|
153
|
+
.env.dev
|
|
153
154
|
.envrc
|
|
154
155
|
.venv
|
|
155
156
|
.autoinference/
|
|
@@ -241,6 +242,22 @@ Network Trash Folder
|
|
|
241
242
|
Temporary Items
|
|
242
243
|
.apdisk
|
|
243
244
|
|
|
244
|
-
# Ephemeral files from
|
|
245
|
+
# Ephemeral files from benchmark runs
|
|
245
246
|
benchmark_results/
|
|
246
247
|
src/autoinference/deployments/
|
|
248
|
+
|
|
249
|
+
# Modal Skills
|
|
250
|
+
.agents/skills/modal/
|
|
251
|
+
.agents/skills/.modal-skill-*
|
|
252
|
+
|
|
253
|
+
# cook/etc is an extension point: local tools live there untracked.
|
|
254
|
+
packages/cook/src/cook/etc/*
|
|
255
|
+
# Local tools that existed before /etc are explicitly added back
|
|
256
|
+
!packages/cook/src/cook/etc/__init__.py
|
|
257
|
+
!packages/cook/src/cook/etc/tests/
|
|
258
|
+
!packages/cook/src/cook/etc/inspect_hf_model.py
|
|
259
|
+
!packages/cook/src/cook/etc/download.py
|
|
260
|
+
!packages/cook/src/cook/etc/quantize.py
|
|
261
|
+
!packages/cook/src/cook/etc/graft_mtp.py
|
|
262
|
+
!packages/cook/src/cook/etc/tune_moe.py
|
|
263
|
+
!packages/cook/src/cook/etc/tool_calling.py
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
[project]
|
|
2
2
|
name = "autoinference-utils"
|
|
3
|
-
version = "0.2.
|
|
3
|
+
version = "0.2.8"
|
|
4
4
|
description = "Shared endpoint abstractions for autoinference deployments"
|
|
5
5
|
readme = "README.md"
|
|
6
6
|
requires-python = ">=3.10"
|
|
@@ -10,7 +10,7 @@ requires = ["hatchling"]
|
|
|
10
10
|
build-backend = "hatchling.build"
|
|
11
11
|
|
|
12
12
|
[tool.hatch.build.targets.wheel]
|
|
13
|
-
packages = ["autoinference_utils"]
|
|
13
|
+
packages = ["src/autoinference_utils"]
|
|
14
14
|
|
|
15
15
|
[tool.hatch.build.targets.sdist]
|
|
16
|
-
include = ["
|
|
16
|
+
include = ["src", "README.md", "pyproject.toml"]
|
|
@@ -33,6 +33,14 @@ MODAL_SESSION_ID_HEADER = "modal-session-id"
|
|
|
33
33
|
ENDPOINTS_REQUIRING_OAI_STREAM_COMPAT: tuple[str, ...] = ("TRTLLMEndpoint",)
|
|
34
34
|
|
|
35
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
|
+
|
|
36
44
|
def server_arg_tokens(flag: str, value: str) -> list[str]:
|
|
37
45
|
"""Render one server arg as argv tokens.
|
|
38
46
|
|
|
@@ -150,7 +158,9 @@ 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
166
|
self.model_path = model_path
|
|
@@ -277,7 +287,7 @@ class SGLangEndpoint(Endpoint):
|
|
|
277
287
|
return cmd
|
|
278
288
|
|
|
279
289
|
def health_check(self) -> str | None:
|
|
280
|
-
url = f"
|
|
290
|
+
url = f"{self.base_url}/health"
|
|
281
291
|
return _health_check(
|
|
282
292
|
url,
|
|
283
293
|
request_timeout=self.health_request_timeout,
|
|
@@ -291,6 +301,7 @@ class SGLangEndpoint(Endpoint):
|
|
|
291
301
|
wait_ready(
|
|
292
302
|
self._proc,
|
|
293
303
|
port=self.worker_port,
|
|
304
|
+
base_url=self.base_url,
|
|
294
305
|
timeout=self.health_timeout,
|
|
295
306
|
poll_interval=self.health_poll_interval,
|
|
296
307
|
request_timeout=self.health_request_timeout,
|
|
@@ -299,6 +310,7 @@ class SGLangEndpoint(Endpoint):
|
|
|
299
310
|
self._bench_server = start_bench_proxy(
|
|
300
311
|
listen_port=self.listen_port,
|
|
301
312
|
upstream_port=self.worker_port,
|
|
313
|
+
upstream_base_url=self.base_url,
|
|
302
314
|
)
|
|
303
315
|
|
|
304
316
|
def stop(self):
|
|
@@ -462,13 +474,16 @@ class RouterEndpoint(Endpoint):
|
|
|
462
474
|
api_key: Optional[str] = None,
|
|
463
475
|
health_timeout: float = 10 * 60,
|
|
464
476
|
health_poll_interval: float = 5.0,
|
|
477
|
+
health_path: str = "/health",
|
|
478
|
+
host: str = "0.0.0.0",
|
|
465
479
|
):
|
|
466
480
|
self.bench_mode = os.environ.get(BENCH_MODE_ENV) == "1"
|
|
467
481
|
self.listen_port = router_port
|
|
468
482
|
actual_router_port = (
|
|
469
483
|
router_port + BENCH_MODE_PORT_OFFSET if self.bench_mode else router_port
|
|
470
484
|
)
|
|
471
|
-
|
|
485
|
+
self.host = host
|
|
486
|
+
super().__init__(base_url=_local_base_url(host, actual_router_port))
|
|
472
487
|
self.pd_config = list(pd_config)
|
|
473
488
|
self.worker_port = (
|
|
474
489
|
worker_port + BENCH_MODE_PORT_OFFSET if self.bench_mode else worker_port
|
|
@@ -478,13 +493,14 @@ class RouterEndpoint(Endpoint):
|
|
|
478
493
|
self.api_key = api_key
|
|
479
494
|
self.health_timeout = health_timeout
|
|
480
495
|
self.health_poll_interval = health_poll_interval
|
|
496
|
+
self.health_path = health_path
|
|
481
497
|
self._proc: Optional[subprocess.Popen] = None
|
|
482
498
|
self._bench_server: Optional[http.server.ThreadingHTTPServer] = None
|
|
483
499
|
|
|
484
500
|
def _build_cmd(self) -> list[str]:
|
|
485
501
|
cmd = [
|
|
486
502
|
"python", "-m", "sglang_router.launch_router",
|
|
487
|
-
"--host", "
|
|
503
|
+
"--host", f"[{self.host}]" if ":" in self.host else self.host,
|
|
488
504
|
"--port", str(self.router_port),
|
|
489
505
|
"--prefill-policy", "cache_aware",
|
|
490
506
|
"--decode-policy", "round_robin",
|
|
@@ -504,7 +520,7 @@ class RouterEndpoint(Endpoint):
|
|
|
504
520
|
cmd.append("--pd-disaggregation")
|
|
505
521
|
|
|
506
522
|
for role, node_ip in self.pd_config:
|
|
507
|
-
node_url =
|
|
523
|
+
node_url = _url(node_ip, self.worker_port)
|
|
508
524
|
if role == "prefill":
|
|
509
525
|
cmd.extend(
|
|
510
526
|
["--prefill", node_url, str(self.prefill_bootstrap_port)]
|
|
@@ -521,7 +537,7 @@ class RouterEndpoint(Endpoint):
|
|
|
521
537
|
def start(self):
|
|
522
538
|
for _, node_ip in self.pd_config:
|
|
523
539
|
_wait_ready_url(
|
|
524
|
-
f"
|
|
540
|
+
f"{_url(node_ip, self.worker_port)}/health",
|
|
525
541
|
timeout=self.health_timeout,
|
|
526
542
|
poll_interval=self.health_poll_interval,
|
|
527
543
|
)
|
|
@@ -529,17 +545,25 @@ class RouterEndpoint(Endpoint):
|
|
|
529
545
|
cmd = self._build_cmd()
|
|
530
546
|
print(f"[router] starting: {shlex.join(cmd)}")
|
|
531
547
|
self._proc = subprocess.Popen(cmd)
|
|
532
|
-
|
|
533
|
-
f"http://localhost:{self.router_port}/health",
|
|
534
|
-
timeout=self.health_timeout,
|
|
535
|
-
poll_interval=self.health_poll_interval,
|
|
536
|
-
)
|
|
548
|
+
self.wait_ready()
|
|
537
549
|
if self.bench_mode:
|
|
538
550
|
self._bench_server = start_bench_proxy(
|
|
539
551
|
listen_port=self.listen_port,
|
|
540
552
|
upstream_port=self.router_port,
|
|
553
|
+
upstream_base_url=self.base_url,
|
|
541
554
|
)
|
|
542
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
|
+
|
|
543
567
|
def stop(self):
|
|
544
568
|
if self._bench_server is not None:
|
|
545
569
|
self._bench_server.shutdown()
|
|
@@ -556,11 +580,12 @@ def start_bench_proxy(
|
|
|
556
580
|
*,
|
|
557
581
|
listen_port: int,
|
|
558
582
|
upstream_port: int,
|
|
583
|
+
upstream_base_url: str | None = None,
|
|
559
584
|
) -> http.server.ThreadingHTTPServer:
|
|
560
585
|
handler_cls = type(
|
|
561
586
|
"BenchProxyHandler",
|
|
562
587
|
(_BenchProxyHandler,),
|
|
563
|
-
{"upstream_port": upstream_port},
|
|
588
|
+
{"upstream_port": upstream_port, "upstream_base_url": upstream_base_url},
|
|
564
589
|
)
|
|
565
590
|
server = http.server.ThreadingHTTPServer(
|
|
566
591
|
("0.0.0.0", listen_port), handler_cls
|
|
@@ -578,6 +603,7 @@ def start_bench_proxy(
|
|
|
578
603
|
|
|
579
604
|
class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
|
|
580
605
|
upstream_port: int = 8000
|
|
606
|
+
upstream_base_url: str | None = None
|
|
581
607
|
|
|
582
608
|
def log_message(self, format, *args):
|
|
583
609
|
return
|
|
@@ -609,6 +635,7 @@ class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
|
|
|
609
635
|
args = payload.get("args") or []
|
|
610
636
|
target = (
|
|
611
637
|
payload.get("target")
|
|
638
|
+
or self.upstream_base_url
|
|
612
639
|
or f"http://localhost:{self.upstream_port}"
|
|
613
640
|
)
|
|
614
641
|
output_dir = payload.get("output_dir") or "/tmp/bench-output"
|
|
@@ -627,7 +654,8 @@ class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
|
|
|
627
654
|
self._send_raw(status, "application/json", body)
|
|
628
655
|
|
|
629
656
|
def _proxy(self, method: str):
|
|
630
|
-
|
|
657
|
+
base_url = self.upstream_base_url or f"http://localhost:{self.upstream_port}"
|
|
658
|
+
url = f"{base_url}{self.path}"
|
|
631
659
|
length = int(self.headers.get("Content-Length", 0))
|
|
632
660
|
body = self.rfile.read(length) if length else None
|
|
633
661
|
forward_headers = {
|
|
@@ -749,12 +777,13 @@ def wait_ready(
|
|
|
749
777
|
port: int,
|
|
750
778
|
timeout: float,
|
|
751
779
|
health_path: str = "/health",
|
|
780
|
+
base_url: str | None = None,
|
|
752
781
|
poll_interval: float = 5.0,
|
|
753
782
|
request_timeout: float = 5.0,
|
|
754
783
|
) -> None:
|
|
755
784
|
"""Poll an HTTP health endpoint until ready, raising if the process dies."""
|
|
756
785
|
deadline = time.time() + timeout
|
|
757
|
-
url = f"
|
|
786
|
+
url = f"{base_url or _local_base_url('127.0.0.1', port)}{health_path}"
|
|
758
787
|
last_error = "no response yet"
|
|
759
788
|
|
|
760
789
|
while time.time() < deadline:
|
|
@@ -788,6 +817,7 @@ def wait_ready(
|
|
|
788
817
|
def warmup_chat_completions(
|
|
789
818
|
*,
|
|
790
819
|
port: int,
|
|
820
|
+
base_url: str | None = None,
|
|
791
821
|
payload: Mapping[str, Any],
|
|
792
822
|
headers: Mapping[str, str] | None = None,
|
|
793
823
|
successful_requests: int = 3,
|
|
@@ -796,7 +826,7 @@ def warmup_chat_completions(
|
|
|
796
826
|
retry_delay: float = 1.0,
|
|
797
827
|
) -> None:
|
|
798
828
|
"""Warm the OpenAI chat completions endpoint with strict retries."""
|
|
799
|
-
url = f"
|
|
829
|
+
url = f"{base_url or _local_base_url('127.0.0.1', port)}/v1/chat/completions"
|
|
800
830
|
request_headers = {"Content-Type": "application/json"}
|
|
801
831
|
if headers:
|
|
802
832
|
request_headers.update(headers)
|
|
@@ -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
|
+
)
|
|
File without changes
|
|
File without changes
|