autoinference-utils 0.2.9__tar.gz → 0.2.10__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {autoinference_utils-0.2.9 → autoinference_utils-0.2.10}/PKG-INFO +1 -1
- {autoinference_utils-0.2.9 → autoinference_utils-0.2.10}/pyproject.toml +1 -1
- {autoinference_utils-0.2.9 → autoinference_utils-0.2.10}/src/autoinference_utils/pd.py +75 -3
- {autoinference_utils-0.2.9 → autoinference_utils-0.2.10}/.gitignore +0 -0
- {autoinference_utils-0.2.9 → autoinference_utils-0.2.10}/README.md +0 -0
- {autoinference_utils-0.2.9 → autoinference_utils-0.2.10}/src/autoinference_utils/__init__.py +0 -0
- {autoinference_utils-0.2.9 → autoinference_utils-0.2.10}/src/autoinference_utils/endpoint.py +0 -0
- {autoinference_utils-0.2.9 → autoinference_utils-0.2.10}/src/autoinference_utils/router.py +0 -0
|
@@ -5,6 +5,8 @@ from __future__ import annotations
|
|
|
5
5
|
import asyncio
|
|
6
6
|
import math
|
|
7
7
|
import os
|
|
8
|
+
import socket
|
|
9
|
+
import struct
|
|
8
10
|
import subprocess
|
|
9
11
|
import threading
|
|
10
12
|
import time
|
|
@@ -40,6 +42,9 @@ _MANAGED_ROUTER_FLAGS = frozenset(
|
|
|
40
42
|
}
|
|
41
43
|
)
|
|
42
44
|
|
|
45
|
+
# Leave headroom inside Modal's 30-second exit-handler deadline.
|
|
46
|
+
_SHUTDOWN_TIMEOUT = 25
|
|
47
|
+
|
|
43
48
|
|
|
44
49
|
def _managed_router_tokens(router_args: Mapping[str, str]) -> list[str]:
|
|
45
50
|
"""Rendered override tokens that would reach a managed flag.
|
|
@@ -177,6 +182,7 @@ class PDEndpoint(Endpoint):
|
|
|
177
182
|
health_failure_timeout: float = 30,
|
|
178
183
|
router_args: Mapping[str, str] | None = None,
|
|
179
184
|
allow_custom_router: bool = False,
|
|
185
|
+
shutdown_port: int | None = None,
|
|
180
186
|
):
|
|
181
187
|
from modal.experimental import get_cluster_info
|
|
182
188
|
|
|
@@ -217,6 +223,13 @@ class PDEndpoint(Endpoint):
|
|
|
217
223
|
f"P/D ratio {ratio} requires a {len(roles)}-container Modal cluster"
|
|
218
224
|
)
|
|
219
225
|
self._host_ip = hosts[cluster.rank]
|
|
226
|
+
self._hosts = hosts
|
|
227
|
+
self._rank = cluster.rank
|
|
228
|
+
self._shutdown_port = (
|
|
229
|
+
router_port + 1 if shutdown_port is None else shutdown_port
|
|
230
|
+
)
|
|
231
|
+
if not 1 <= self._shutdown_port <= 65535:
|
|
232
|
+
raise ValueError("shutdown_port must be between 1 and 65535")
|
|
220
233
|
self.role = roles[cluster.rank]
|
|
221
234
|
super().__init__(_url(hosts[0], router_port))
|
|
222
235
|
self.engine = engine
|
|
@@ -266,6 +279,8 @@ class PDEndpoint(Endpoint):
|
|
|
266
279
|
self._stopped = threading.Event()
|
|
267
280
|
self._lock = threading.Lock()
|
|
268
281
|
self._started = False
|
|
282
|
+
self._serving = False
|
|
283
|
+
self._shutdown_listener: socket.socket | None = None
|
|
269
284
|
self._threads: list[threading.Thread] = []
|
|
270
285
|
|
|
271
286
|
def start(self, *, warmup: Callable[[], None] | None = None) -> None:
|
|
@@ -280,6 +295,13 @@ class PDEndpoint(Endpoint):
|
|
|
280
295
|
SGLANG_ENABLE_HEALTH_ENDPOINT_GENERATION="1",
|
|
281
296
|
)
|
|
282
297
|
try:
|
|
298
|
+
if self._rank == 0:
|
|
299
|
+
self._shutdown_listener = socket.socket(socket.AF_INET6)
|
|
300
|
+
self._shutdown_listener.setsockopt(
|
|
301
|
+
socket.SOL_SOCKET, socket.SO_REUSEADDR, 1
|
|
302
|
+
)
|
|
303
|
+
self._shutdown_listener.bind(("::", self._shutdown_port))
|
|
304
|
+
self._shutdown_listener.listen(len(self._hosts) - 1)
|
|
283
305
|
for endpoint in (self.engine, self.router):
|
|
284
306
|
if endpoint is not None:
|
|
285
307
|
endpoint.start()
|
|
@@ -291,6 +313,7 @@ class PDEndpoint(Endpoint):
|
|
|
291
313
|
for _, host in self.router.pd_config:
|
|
292
314
|
self._watch_health(_url(host, self.engine.worker_port) + "/health")
|
|
293
315
|
self._watch_health(self.router.base_url + "/health")
|
|
316
|
+
self._serving = True
|
|
294
317
|
except BaseException:
|
|
295
318
|
self.stop()
|
|
296
319
|
raise
|
|
@@ -352,7 +375,7 @@ class PDEndpoint(Endpoint):
|
|
|
352
375
|
if self._stopped.is_set():
|
|
353
376
|
return
|
|
354
377
|
self._stopped.set()
|
|
355
|
-
deadline = time.monotonic() +
|
|
378
|
+
deadline = time.monotonic() + _SHUTDOWN_TIMEOUT
|
|
356
379
|
try:
|
|
357
380
|
for endpoint in (self.router, self.engine):
|
|
358
381
|
if endpoint is None:
|
|
@@ -364,9 +387,58 @@ class PDEndpoint(Endpoint):
|
|
|
364
387
|
process.wait(timeout=max(0, deadline - time.monotonic()))
|
|
365
388
|
except subprocess.TimeoutExpired:
|
|
366
389
|
process.kill()
|
|
390
|
+
process.wait()
|
|
367
391
|
endpoint.stop()
|
|
368
392
|
finally:
|
|
369
393
|
for thread in self._threads:
|
|
370
|
-
thread.join(
|
|
371
|
-
|
|
394
|
+
thread.join(timeout=max(0, deadline - time.monotonic()))
|
|
395
|
+
if self._serving:
|
|
396
|
+
if self._rank == 0:
|
|
397
|
+
self._wait_for_followers(deadline)
|
|
398
|
+
else:
|
|
399
|
+
self._notify_leader(deadline)
|
|
400
|
+
if self._shutdown_listener is not None:
|
|
401
|
+
self._shutdown_listener.close()
|
|
402
|
+
self._shutdown_listener = None
|
|
403
|
+
|
|
404
|
+
def _notify_leader(self, deadline: float) -> None:
|
|
405
|
+
while True:
|
|
406
|
+
remaining = deadline - time.monotonic()
|
|
407
|
+
if remaining <= 0:
|
|
408
|
+
print("[pd] Failed to notify rank zero of shutdown", flush=True)
|
|
409
|
+
return
|
|
410
|
+
try:
|
|
411
|
+
with socket.create_connection(
|
|
412
|
+
(self._hosts[0], self._shutdown_port), timeout=min(remaining, 1)
|
|
413
|
+
) as connection:
|
|
414
|
+
connection.sendall(struct.pack("!I", self._rank))
|
|
415
|
+
return
|
|
416
|
+
except OSError:
|
|
417
|
+
time.sleep(min(0.1, max(0, deadline - time.monotonic())))
|
|
418
|
+
|
|
419
|
+
def _wait_for_followers(self, deadline: float) -> None:
|
|
420
|
+
listener = self._shutdown_listener
|
|
421
|
+
if listener is None:
|
|
422
|
+
return
|
|
423
|
+
pending = set(range(1, len(self._hosts)))
|
|
424
|
+
while pending:
|
|
425
|
+
remaining = deadline - time.monotonic()
|
|
426
|
+
if remaining <= 0:
|
|
427
|
+
print(
|
|
428
|
+
f"[pd] Timed out waiting for ranks {sorted(pending)} to stop",
|
|
429
|
+
flush=True,
|
|
372
430
|
)
|
|
431
|
+
return
|
|
432
|
+
listener.settimeout(min(remaining, 1))
|
|
433
|
+
try:
|
|
434
|
+
connection, _ = listener.accept()
|
|
435
|
+
except TimeoutError:
|
|
436
|
+
continue
|
|
437
|
+
with connection:
|
|
438
|
+
connection.settimeout(min(remaining, 1))
|
|
439
|
+
try:
|
|
440
|
+
payload = connection.recv(4, socket.MSG_WAITALL)
|
|
441
|
+
except TimeoutError:
|
|
442
|
+
continue
|
|
443
|
+
if len(payload) == 4:
|
|
444
|
+
pending.discard(struct.unpack("!I", payload)[0])
|
|
File without changes
|
|
File without changes
|
{autoinference_utils-0.2.9 → autoinference_utils-0.2.10}/src/autoinference_utils/__init__.py
RENAMED
|
File without changes
|
{autoinference_utils-0.2.9 → autoinference_utils-0.2.10}/src/autoinference_utils/endpoint.py
RENAMED
|
File without changes
|
|
File without changes
|