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.
@@ -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 Experiments
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
  Metadata-Version: 2.5
2
2
  Name: autoinference-utils
3
- Version: 0.2.6
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
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "autoinference-utils"
3
- version = "0.2.6"
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 = ["autoinference_utils", "README.md", "pyproject.toml"]
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
- super().__init__(base_url=f"http://localhost:{sglang_port}")
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"http://127.0.0.1:{self.worker_port}/health"
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
- super().__init__(base_url=f"http://localhost:{actual_router_port}")
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", "0.0.0.0",
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 = f"http://{node_ip}:{self.worker_port}"
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"http://{node_ip}:{self.worker_port}/health",
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
- _wait_ready_url(
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
- url = f"http://localhost:{self.upstream_port}{self.path}"
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"http://127.0.0.1:{port}{health_path}"
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"http://127.0.0.1:{port}/v1/chat/completions"
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
+ )