autoinference-utils 0.2.8__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.8 → autoinference_utils-0.2.10}/PKG-INFO +1 -1
- {autoinference_utils-0.2.8 → autoinference_utils-0.2.10}/pyproject.toml +1 -1
- {autoinference_utils-0.2.8 → autoinference_utils-0.2.10}/src/autoinference_utils/endpoint.py +94 -104
- {autoinference_utils-0.2.8 → autoinference_utils-0.2.10}/src/autoinference_utils/pd.py +143 -14
- autoinference_utils-0.2.10/src/autoinference_utils/router.py +103 -0
- {autoinference_utils-0.2.8 → autoinference_utils-0.2.10}/.gitignore +0 -0
- {autoinference_utils-0.2.8 → autoinference_utils-0.2.10}/README.md +0 -0
- {autoinference_utils-0.2.8 → autoinference_utils-0.2.10}/src/autoinference_utils/__init__.py +0 -0
{autoinference_utils-0.2.8 → autoinference_utils-0.2.10}/src/autoinference_utils/endpoint.py
RENAMED
|
@@ -27,7 +27,7 @@ BENCH_MODE_ENV = "ENDPOINT_BENCH_MODE"
|
|
|
27
27
|
DUMMY_WEIGHTS_ENV = "ENDPOINT_DUMMY"
|
|
28
28
|
BENCH_MODE_PORT_OFFSET = 10000
|
|
29
29
|
|
|
30
|
-
MODAL_FLASH_REQUEST_UUID_HEADER= "modal-flash-request-uuid"
|
|
30
|
+
MODAL_FLASH_REQUEST_UUID_HEADER = "modal-flash-request-uuid"
|
|
31
31
|
MODAL_SESSION_ID_HEADER = "modal-session-id"
|
|
32
32
|
|
|
33
33
|
ENDPOINTS_REQUIRING_OAI_STREAM_COMPAT: tuple[str, ...] = ("TRTLLMEndpoint",)
|
|
@@ -63,9 +63,7 @@ def _flag_name(key: str) -> str:
|
|
|
63
63
|
return key.partition("=")[0]
|
|
64
64
|
|
|
65
65
|
|
|
66
|
-
def get_server_arg(
|
|
67
|
-
server_args: Mapping[str, str], flag: str
|
|
68
|
-
) -> tuple[str, str] | None:
|
|
66
|
+
def get_server_arg(server_args: Mapping[str, str], flag: str) -> tuple[str, str] | None:
|
|
69
67
|
"""Return the key and value setting a flag under either spelling."""
|
|
70
68
|
for key, value in server_args.items():
|
|
71
69
|
if _flag_name(key) == flag:
|
|
@@ -176,9 +174,7 @@ class SGLangEndpoint(Endpoint):
|
|
|
176
174
|
self.disaggregation_mode = disaggregation_mode
|
|
177
175
|
self.prefill_bootstrap_port = prefill_bootstrap_port
|
|
178
176
|
self.launcher_module = launcher_module
|
|
179
|
-
self.extra_server_args = (
|
|
180
|
-
dict(extra_server_args) if extra_server_args else {}
|
|
181
|
-
)
|
|
177
|
+
self.extra_server_args = dict(extra_server_args) if extra_server_args else {}
|
|
182
178
|
self.health_timeout = health_timeout
|
|
183
179
|
self.health_poll_interval = health_poll_interval
|
|
184
180
|
self.health_request_timeout = health_request_timeout
|
|
@@ -193,17 +189,16 @@ class SGLangEndpoint(Endpoint):
|
|
|
193
189
|
self._bench_server: Optional[http.server.ThreadingHTTPServer] = None
|
|
194
190
|
|
|
195
191
|
if self.disaggregation_mode not in (None, "prefill", "decode"):
|
|
196
|
-
raise ValueError(
|
|
197
|
-
"disaggregation_mode must be None, 'prefill', or 'decode'"
|
|
198
|
-
)
|
|
192
|
+
raise ValueError("disaggregation_mode must be None, 'prefill', or 'decode'")
|
|
199
193
|
|
|
200
194
|
if os.environ.get(DUMMY_WEIGHTS_ENV) == "1":
|
|
201
|
-
if (
|
|
202
|
-
self.
|
|
203
|
-
and not has_server_arg(self.extra_server_args, "--load-format")
|
|
195
|
+
if self.load_format is None and not has_server_arg(
|
|
196
|
+
self.extra_server_args, "--load-format"
|
|
204
197
|
):
|
|
205
198
|
self.load_format = "dummy"
|
|
206
|
-
if not has_server_arg(
|
|
199
|
+
if not has_server_arg(
|
|
200
|
+
self.extra_server_args, "--model-loader-extra-config"
|
|
201
|
+
):
|
|
207
202
|
self.extra_server_args["--model-loader-extra-config"] = "{}"
|
|
208
203
|
|
|
209
204
|
# Request logging enabled, log only metadata by default
|
|
@@ -214,7 +209,9 @@ class SGLangEndpoint(Endpoint):
|
|
|
214
209
|
# Set request logging level
|
|
215
210
|
level = get_server_arg(self.extra_server_args, "--log-requests-level")
|
|
216
211
|
if level is None:
|
|
217
|
-
self.extra_server_args["--log-requests-level"] = str(
|
|
212
|
+
self.extra_server_args["--log-requests-level"] = str(
|
|
213
|
+
self.log_requests_level
|
|
214
|
+
)
|
|
218
215
|
else:
|
|
219
216
|
key, value = level
|
|
220
217
|
print(
|
|
@@ -236,19 +233,21 @@ class SGLangEndpoint(Endpoint):
|
|
|
236
233
|
self.log_request_headers = ",".join(sglang_log_request_headers)
|
|
237
234
|
os.environ["SGLANG_LOG_REQUEST_HEADERS"] = self.log_request_headers
|
|
238
235
|
|
|
239
|
-
|
|
240
236
|
def _build_cmd(self) -> list[str]:
|
|
241
237
|
cmd = [
|
|
242
|
-
"python",
|
|
243
|
-
"
|
|
244
|
-
|
|
245
|
-
"--
|
|
238
|
+
"python",
|
|
239
|
+
"-m",
|
|
240
|
+
self.launcher_module,
|
|
241
|
+
"--host",
|
|
242
|
+
"0.0.0.0",
|
|
243
|
+
"--port",
|
|
244
|
+
str(self.worker_port),
|
|
245
|
+
"--model-path",
|
|
246
|
+
self.model_path,
|
|
246
247
|
]
|
|
247
248
|
|
|
248
249
|
if self.speculative_model_path is not None:
|
|
249
|
-
cmd.extend(
|
|
250
|
-
["--speculative-draft-model-path", self.speculative_model_path]
|
|
251
|
-
)
|
|
250
|
+
cmd.extend(["--speculative-draft-model-path", self.speculative_model_path])
|
|
252
251
|
if self.load_format is not None:
|
|
253
252
|
cmd.extend(["--load-format", self.load_format])
|
|
254
253
|
if self.tp is not None:
|
|
@@ -273,14 +272,18 @@ class SGLangEndpoint(Endpoint):
|
|
|
273
272
|
raise ValueError("dist_init_host is required when nnodes > 1")
|
|
274
273
|
cmd.extend(
|
|
275
274
|
[
|
|
276
|
-
"--nnodes",
|
|
277
|
-
|
|
275
|
+
"--nnodes",
|
|
276
|
+
str(self.nnodes),
|
|
277
|
+
"--node-rank",
|
|
278
|
+
str(self.node_rank),
|
|
278
279
|
"--dist-init-addr",
|
|
279
280
|
f"{self.dist_init_host}:{self.dist_init_port}",
|
|
280
281
|
]
|
|
281
282
|
)
|
|
282
283
|
|
|
283
|
-
merged = _merge_server_args(
|
|
284
|
+
merged = _merge_server_args(
|
|
285
|
+
self.DEFAULT_OPERATIONAL_ARGS, self.extra_server_args
|
|
286
|
+
)
|
|
284
287
|
for key, value in merged.items():
|
|
285
288
|
cmd.extend(server_arg_tokens(key, value))
|
|
286
289
|
|
|
@@ -345,9 +348,7 @@ class VLLMEndpoint(Endpoint):
|
|
|
345
348
|
super().__init__(base_url=f"http://localhost:{vllm_port}")
|
|
346
349
|
self.model = model
|
|
347
350
|
self.worker_port = vllm_port
|
|
348
|
-
self.extra_server_args = (
|
|
349
|
-
dict(extra_server_args) if extra_server_args else {}
|
|
350
|
-
)
|
|
351
|
+
self.extra_server_args = dict(extra_server_args) if extra_server_args else {}
|
|
351
352
|
self.health_timeout = health_timeout
|
|
352
353
|
self.health_poll_interval = health_poll_interval
|
|
353
354
|
self.health_request_timeout = health_request_timeout
|
|
@@ -360,10 +361,15 @@ class VLLMEndpoint(Endpoint):
|
|
|
360
361
|
|
|
361
362
|
def _build_cmd(self) -> list[str]:
|
|
362
363
|
cmd = [
|
|
363
|
-
"python",
|
|
364
|
-
"
|
|
365
|
-
"
|
|
366
|
-
"--
|
|
364
|
+
"python",
|
|
365
|
+
"-m",
|
|
366
|
+
"vllm.entrypoints.openai.api_server",
|
|
367
|
+
"--host",
|
|
368
|
+
"0.0.0.0",
|
|
369
|
+
"--port",
|
|
370
|
+
str(self.worker_port),
|
|
371
|
+
"--model",
|
|
372
|
+
self.model,
|
|
367
373
|
]
|
|
368
374
|
for key, value in self.extra_server_args.items():
|
|
369
375
|
cmd.extend(server_arg_tokens(key, value))
|
|
@@ -417,9 +423,7 @@ class TRTLLMEndpoint(Endpoint):
|
|
|
417
423
|
super().__init__(base_url=f"http://localhost:{worker_port}")
|
|
418
424
|
self.model = model
|
|
419
425
|
self.worker_port = worker_port
|
|
420
|
-
self.extra_server_args = (
|
|
421
|
-
dict(extra_server_args) if extra_server_args else {}
|
|
422
|
-
)
|
|
426
|
+
self.extra_server_args = dict(extra_server_args) if extra_server_args else {}
|
|
423
427
|
self.health_timeout = health_timeout
|
|
424
428
|
self.health_poll_interval = health_poll_interval
|
|
425
429
|
self.health_request_timeout = health_request_timeout
|
|
@@ -429,8 +433,10 @@ class TRTLLMEndpoint(Endpoint):
|
|
|
429
433
|
cmd = [
|
|
430
434
|
"trtllm-serve",
|
|
431
435
|
self.model,
|
|
432
|
-
"--host",
|
|
433
|
-
"
|
|
436
|
+
"--host",
|
|
437
|
+
"0.0.0.0",
|
|
438
|
+
"--port",
|
|
439
|
+
str(self.worker_port),
|
|
434
440
|
]
|
|
435
441
|
for key, value in self.extra_server_args.items():
|
|
436
442
|
cmd.extend(server_arg_tokens(key, value))
|
|
@@ -499,18 +505,30 @@ class RouterEndpoint(Endpoint):
|
|
|
499
505
|
|
|
500
506
|
def _build_cmd(self) -> list[str]:
|
|
501
507
|
cmd = [
|
|
502
|
-
"python",
|
|
503
|
-
"
|
|
504
|
-
"
|
|
505
|
-
"--
|
|
506
|
-
"
|
|
507
|
-
"--
|
|
508
|
-
|
|
509
|
-
"--
|
|
510
|
-
"
|
|
511
|
-
"--
|
|
508
|
+
"python",
|
|
509
|
+
"-m",
|
|
510
|
+
"sglang_router.launch_router",
|
|
511
|
+
"--host",
|
|
512
|
+
f"[{self.host}]" if ":" in self.host else self.host,
|
|
513
|
+
"--port",
|
|
514
|
+
str(self.router_port),
|
|
515
|
+
"--prefill-policy",
|
|
516
|
+
"cache_aware",
|
|
517
|
+
"--decode-policy",
|
|
518
|
+
"round_robin",
|
|
519
|
+
"--max-concurrent-requests",
|
|
520
|
+
"128",
|
|
521
|
+
"--rate-limit-tokens-per-second",
|
|
522
|
+
"0",
|
|
523
|
+
"--queue-size",
|
|
524
|
+
"0",
|
|
525
|
+
"--health-check-timeout-secs",
|
|
526
|
+
"600",
|
|
527
|
+
"--log-level",
|
|
528
|
+
"info",
|
|
512
529
|
"--disable-circuit-breaker",
|
|
513
|
-
"--request-timeout-secs",
|
|
530
|
+
"--request-timeout-secs",
|
|
531
|
+
"3600",
|
|
514
532
|
]
|
|
515
533
|
|
|
516
534
|
if self.api_key is not None:
|
|
@@ -522,9 +540,7 @@ class RouterEndpoint(Endpoint):
|
|
|
522
540
|
for role, node_ip in self.pd_config:
|
|
523
541
|
node_url = _url(node_ip, self.worker_port)
|
|
524
542
|
if role == "prefill":
|
|
525
|
-
cmd.extend(
|
|
526
|
-
["--prefill", node_url, str(self.prefill_bootstrap_port)]
|
|
527
|
-
)
|
|
543
|
+
cmd.extend(["--prefill", node_url, str(self.prefill_bootstrap_port)])
|
|
528
544
|
elif role == "decode":
|
|
529
545
|
cmd.extend(["--decode", node_url])
|
|
530
546
|
elif role == "worker":
|
|
@@ -576,6 +592,7 @@ class RouterEndpoint(Endpoint):
|
|
|
576
592
|
# and this proxy listens on worker_port. POST /bench shells out to a benchmark
|
|
577
593
|
# task via run_bench(); every other path is forwarded to the upstream server.
|
|
578
594
|
|
|
595
|
+
|
|
579
596
|
def start_bench_proxy(
|
|
580
597
|
*,
|
|
581
598
|
listen_port: int,
|
|
@@ -587,17 +604,12 @@ def start_bench_proxy(
|
|
|
587
604
|
(_BenchProxyHandler,),
|
|
588
605
|
{"upstream_port": upstream_port, "upstream_base_url": upstream_base_url},
|
|
589
606
|
)
|
|
590
|
-
server = http.server.ThreadingHTTPServer(
|
|
591
|
-
("0.0.0.0", listen_port), handler_cls
|
|
592
|
-
)
|
|
607
|
+
server = http.server.ThreadingHTTPServer(("0.0.0.0", listen_port), handler_cls)
|
|
593
608
|
thread = threading.Thread(
|
|
594
609
|
target=server.serve_forever, daemon=True, name="bench-proxy"
|
|
595
610
|
)
|
|
596
611
|
thread.start()
|
|
597
|
-
print(
|
|
598
|
-
f"[bench-proxy] listening on :{listen_port}, "
|
|
599
|
-
f"forwarding to :{upstream_port}"
|
|
600
|
-
)
|
|
612
|
+
print(f"[bench-proxy] listening on :{listen_port}, forwarding to :{upstream_port}")
|
|
601
613
|
return server
|
|
602
614
|
|
|
603
615
|
|
|
@@ -627,9 +639,7 @@ class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
|
|
|
627
639
|
try:
|
|
628
640
|
payload = self._read_json_body()
|
|
629
641
|
except (ValueError, json.JSONDecodeError) as exc:
|
|
630
|
-
return self._send_json(
|
|
631
|
-
400, {"ok": False, "error": f"bad body: {exc}"}
|
|
632
|
-
)
|
|
642
|
+
return self._send_json(400, {"ok": False, "error": f"bad body: {exc}"})
|
|
633
643
|
|
|
634
644
|
benchmark = payload.get("benchmark") or ""
|
|
635
645
|
args = payload.get("args") or []
|
|
@@ -641,13 +651,9 @@ class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
|
|
|
641
651
|
output_dir = payload.get("output_dir") or "/tmp/bench-output"
|
|
642
652
|
|
|
643
653
|
if not benchmark:
|
|
644
|
-
return self._send_json(
|
|
645
|
-
400, {"ok": False, "error": "benchmark is required"}
|
|
646
|
-
)
|
|
654
|
+
return self._send_json(400, {"ok": False, "error": "benchmark is required"})
|
|
647
655
|
if not isinstance(args, list):
|
|
648
|
-
return self._send_json(
|
|
649
|
-
400, {"ok": False, "error": "args must be a list"}
|
|
650
|
-
)
|
|
656
|
+
return self._send_json(400, {"ok": False, "error": "args must be a list"})
|
|
651
657
|
|
|
652
658
|
rc, body = run_bench(benchmark, list(args), target, output_dir)
|
|
653
659
|
status = 200 if rc == 0 else 500
|
|
@@ -699,9 +705,7 @@ class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
|
|
|
699
705
|
return json.loads(raw)
|
|
700
706
|
|
|
701
707
|
def _send_json(self, status: int, obj: Mapping[str, Any]):
|
|
702
|
-
self._send_raw(
|
|
703
|
-
status, "application/json", json.dumps(obj).encode()
|
|
704
|
-
)
|
|
708
|
+
self._send_raw(status, "application/json", json.dumps(obj).encode())
|
|
705
709
|
|
|
706
710
|
def _send_raw(self, status: int, content_type: str, body: bytes):
|
|
707
711
|
self.send_response(status)
|
|
@@ -726,18 +730,21 @@ def run_bench(
|
|
|
726
730
|
return 500, json.dumps(
|
|
727
731
|
{
|
|
728
732
|
"ok": False,
|
|
729
|
-
"error": (
|
|
730
|
-
"autoinference.benchmarks not importable in container"
|
|
731
|
-
),
|
|
733
|
+
"error": ("autoinference.benchmarks not importable in container"),
|
|
732
734
|
"volume_path": output_dir,
|
|
733
735
|
}
|
|
734
736
|
).encode()
|
|
735
737
|
|
|
736
738
|
cmd = [
|
|
737
|
-
sys.executable,
|
|
738
|
-
"
|
|
739
|
-
"
|
|
740
|
-
|
|
739
|
+
sys.executable,
|
|
740
|
+
"-m",
|
|
741
|
+
"invoke",
|
|
742
|
+
"--search-root",
|
|
743
|
+
search_root,
|
|
744
|
+
"-c",
|
|
745
|
+
"tasks",
|
|
746
|
+
benchmark,
|
|
747
|
+
*[str(a) for a in args],
|
|
741
748
|
]
|
|
742
749
|
env = {
|
|
743
750
|
**os.environ,
|
|
@@ -788,9 +795,7 @@ def wait_ready(
|
|
|
788
795
|
|
|
789
796
|
while time.time() < deadline:
|
|
790
797
|
try:
|
|
791
|
-
error = _health_check(
|
|
792
|
-
url, request_timeout=request_timeout, process=process
|
|
793
|
-
)
|
|
798
|
+
error = _health_check(url, request_timeout=request_timeout, process=process)
|
|
794
799
|
except subprocess.CalledProcessError as exc:
|
|
795
800
|
print(
|
|
796
801
|
f"[endpoint] !!! server process exited with code "
|
|
@@ -804,13 +809,11 @@ def wait_ready(
|
|
|
804
809
|
time.sleep(poll_interval)
|
|
805
810
|
|
|
806
811
|
print(
|
|
807
|
-
f"[endpoint] !!! health check timed out after {timeout}s; "
|
|
808
|
-
f"last={last_error}",
|
|
812
|
+
f"[endpoint] !!! health check timed out after {timeout}s; last={last_error}",
|
|
809
813
|
flush=True,
|
|
810
814
|
)
|
|
811
815
|
raise TimeoutError(
|
|
812
|
-
f"Health check timed out after {timeout}s for {url}. "
|
|
813
|
-
f"Last error: {last_error}"
|
|
816
|
+
f"Health check timed out after {timeout}s for {url}. Last error: {last_error}"
|
|
814
817
|
)
|
|
815
818
|
|
|
816
819
|
|
|
@@ -839,9 +842,7 @@ def warmup_chat_completions(
|
|
|
839
842
|
request_timeout=request_timeout,
|
|
840
843
|
max_attempts=max_attempts_per_request,
|
|
841
844
|
retry_delay=retry_delay,
|
|
842
|
-
description=(
|
|
843
|
-
f"warmup request {request_idx + 1}/{successful_requests}"
|
|
844
|
-
),
|
|
845
|
+
description=(f"warmup request {request_idx + 1}/{successful_requests}"),
|
|
845
846
|
)
|
|
846
847
|
|
|
847
848
|
|
|
@@ -863,9 +864,7 @@ def validate_embeddings_endpoint(
|
|
|
863
864
|
1 if all(isinstance(item, int) for item in inputs) else len(inputs)
|
|
864
865
|
)
|
|
865
866
|
else:
|
|
866
|
-
raise ValueError(
|
|
867
|
-
"embedding validation payload must contain non-empty input"
|
|
868
|
-
)
|
|
867
|
+
raise ValueError("embedding validation payload must contain non-empty input")
|
|
869
868
|
|
|
870
869
|
request_headers = {"Content-Type": "application/json"}
|
|
871
870
|
if headers:
|
|
@@ -928,10 +927,7 @@ def start_heartbeat_thread(
|
|
|
928
927
|
try:
|
|
929
928
|
error = health_check_fn()
|
|
930
929
|
except subprocess.CalledProcessError as exc:
|
|
931
|
-
print(
|
|
932
|
-
f"[heartbeat] server process exited with code "
|
|
933
|
-
f"{exc.returncode}"
|
|
934
|
-
)
|
|
930
|
+
print(f"[heartbeat] server process exited with code {exc.returncode}")
|
|
935
931
|
on_failure()
|
|
936
932
|
return
|
|
937
933
|
if error is None:
|
|
@@ -943,10 +939,7 @@ def start_heartbeat_thread(
|
|
|
943
939
|
f"({consecutive_failures}/{max_consecutive_failures})"
|
|
944
940
|
)
|
|
945
941
|
if consecutive_failures >= max_consecutive_failures:
|
|
946
|
-
print(
|
|
947
|
-
"[heartbeat] sustained health-check failure, "
|
|
948
|
-
"invoking on_failure"
|
|
949
|
-
)
|
|
942
|
+
print("[heartbeat] sustained health-check failure, invoking on_failure")
|
|
950
943
|
on_failure()
|
|
951
944
|
return
|
|
952
945
|
|
|
@@ -1031,8 +1024,7 @@ def _wait_ready_url(
|
|
|
1031
1024
|
last_error = error
|
|
1032
1025
|
time.sleep(poll_interval)
|
|
1033
1026
|
raise TimeoutError(
|
|
1034
|
-
f"Timed out after {timeout}s waiting for {url}. "
|
|
1035
|
-
f"Last error: {last_error}"
|
|
1027
|
+
f"Timed out after {timeout}s waiting for {url}. Last error: {last_error}"
|
|
1036
1028
|
)
|
|
1037
1029
|
|
|
1038
1030
|
|
|
@@ -1050,9 +1042,7 @@ def _post_json(
|
|
|
1050
1042
|
timeout: float | None = None,
|
|
1051
1043
|
) -> bytes:
|
|
1052
1044
|
body = json.dumps(payload).encode("utf-8")
|
|
1053
|
-
req = urllib.request.Request(
|
|
1054
|
-
url, data=body, headers=dict(headers), method="POST"
|
|
1055
|
-
)
|
|
1045
|
+
req = urllib.request.Request(url, data=body, headers=dict(headers), method="POST")
|
|
1056
1046
|
with urllib.request.urlopen(req, timeout=timeout) as resp:
|
|
1057
1047
|
return resp.read()
|
|
1058
1048
|
|
|
@@ -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
|
|
@@ -19,7 +21,52 @@ from .endpoint import (
|
|
|
19
21
|
_local_base_url,
|
|
20
22
|
_merge_server_args,
|
|
21
23
|
_url,
|
|
24
|
+
server_arg_tokens,
|
|
22
25
|
)
|
|
26
|
+
from .router import verify_router
|
|
27
|
+
|
|
28
|
+
# PDEndpoint owns the topology and every flag it mirrors in its own supervision
|
|
29
|
+
# (drain deadline, health probes); recipes tune everything else via router_args.
|
|
30
|
+
_MANAGED_ROUTER_FLAGS = frozenset(
|
|
31
|
+
{
|
|
32
|
+
"--host",
|
|
33
|
+
"--port",
|
|
34
|
+
"--pd-disaggregation",
|
|
35
|
+
"--prefill",
|
|
36
|
+
"--decode",
|
|
37
|
+
"--api-key",
|
|
38
|
+
"--shutdown-grace-period-secs",
|
|
39
|
+
"--health-check-interval-secs",
|
|
40
|
+
"--health-check-timeout-secs",
|
|
41
|
+
"--health-failure-threshold",
|
|
42
|
+
}
|
|
43
|
+
)
|
|
44
|
+
|
|
45
|
+
# Leave headroom inside Modal's 30-second exit-handler deadline.
|
|
46
|
+
_SHUTDOWN_TIMEOUT = 25
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def _managed_router_tokens(router_args: Mapping[str, str]) -> list[str]:
|
|
50
|
+
"""Rendered override tokens that would reach a managed flag.
|
|
51
|
+
|
|
52
|
+
Checks argv as the router parses it: values split into extra tokens, and
|
|
53
|
+
argparse accepts any unambiguous prefix of a long flag.
|
|
54
|
+
"""
|
|
55
|
+
tokens = [
|
|
56
|
+
token
|
|
57
|
+
for key, value in router_args.items()
|
|
58
|
+
for token in server_arg_tokens(key, value)
|
|
59
|
+
]
|
|
60
|
+
return sorted(
|
|
61
|
+
{
|
|
62
|
+
token
|
|
63
|
+
for token in tokens
|
|
64
|
+
if token.startswith("-")
|
|
65
|
+
and any(
|
|
66
|
+
flag.startswith(_flag_name(token)) for flag in _MANAGED_ROUTER_FLAGS
|
|
67
|
+
)
|
|
68
|
+
}
|
|
69
|
+
)
|
|
23
70
|
|
|
24
71
|
|
|
25
72
|
def _configure_fabric() -> None:
|
|
@@ -61,6 +108,7 @@ class _PDRouter(RouterEndpoint):
|
|
|
61
108
|
health_interval,
|
|
62
109
|
health_request_timeout,
|
|
63
110
|
health_failure_threshold,
|
|
111
|
+
router_args,
|
|
64
112
|
**kwargs,
|
|
65
113
|
):
|
|
66
114
|
super().__init__(health_path="/readiness", host="::", **kwargs)
|
|
@@ -69,6 +117,7 @@ class _PDRouter(RouterEndpoint):
|
|
|
69
117
|
self.health_interval = health_interval
|
|
70
118
|
self.health_request_timeout = health_request_timeout
|
|
71
119
|
self.health_failure_threshold = health_failure_threshold
|
|
120
|
+
self.router_args = dict(router_args)
|
|
72
121
|
|
|
73
122
|
def _build_cmd(self) -> list[str]:
|
|
74
123
|
cmd = [
|
|
@@ -85,20 +134,21 @@ class _PDRouter(RouterEndpoint):
|
|
|
85
134
|
cmd.extend([f"--{role}", _url(host, self.worker_port)])
|
|
86
135
|
if role == "prefill":
|
|
87
136
|
cmd.append(str(self.prefill_bootstrap_port))
|
|
88
|
-
|
|
137
|
+
defaults = {
|
|
89
138
|
"--prefill-policy": "cache_aware",
|
|
90
139
|
"--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,
|
|
140
|
+
"--max-concurrent-requests": str(self.max_concurrent_requests),
|
|
141
|
+
"--queue-size": "0",
|
|
142
|
+
"--rate-limit-tokens-per-second": "0",
|
|
143
|
+
"--health-check-interval-secs": str(self.health_interval),
|
|
144
|
+
"--health-check-timeout-secs": str(self.health_request_timeout),
|
|
145
|
+
"--health-failure-threshold": str(self.health_failure_threshold),
|
|
146
|
+
"--shutdown-grace-period-secs": str(self.drain_timeout),
|
|
147
|
+
"--request-timeout-secs": "1800",
|
|
99
148
|
"--log-level": "info",
|
|
100
|
-
}
|
|
101
|
-
|
|
149
|
+
}
|
|
150
|
+
for flag, value in _merge_server_args(defaults, self.router_args).items():
|
|
151
|
+
cmd.extend(server_arg_tokens(flag, value))
|
|
102
152
|
if self.api_key is not None:
|
|
103
153
|
cmd.extend(["--api-key", self.api_key])
|
|
104
154
|
return cmd
|
|
@@ -110,6 +160,8 @@ class PDEndpoint(Endpoint):
|
|
|
110
160
|
Call start/stop from Modal enter/exit hooks. Model-specific arguments and
|
|
111
161
|
warmup stay in the recipe. The image must support NIXL and graceful shutdown.
|
|
112
162
|
The ratio is (prefill, decode) counts; prefill ranks precede decode ranks.
|
|
163
|
+
The router is the build pinned in autoinference_utils.router; router_args
|
|
164
|
+
override its default flags, and allow_custom_router skips the pin check.
|
|
113
165
|
"""
|
|
114
166
|
|
|
115
167
|
def __init__(
|
|
@@ -128,6 +180,9 @@ class PDEndpoint(Endpoint):
|
|
|
128
180
|
router_health_request_timeout: int = 5,
|
|
129
181
|
health_failure_threshold: int = 3,
|
|
130
182
|
health_failure_timeout: float = 30,
|
|
183
|
+
router_args: Mapping[str, str] | None = None,
|
|
184
|
+
allow_custom_router: bool = False,
|
|
185
|
+
shutdown_port: int | None = None,
|
|
131
186
|
):
|
|
132
187
|
from modal.experimental import get_cluster_info
|
|
133
188
|
|
|
@@ -156,6 +211,10 @@ class PDEndpoint(Endpoint):
|
|
|
156
211
|
raise ValueError(
|
|
157
212
|
"ratio must contain positive integer prefill and decode counts"
|
|
158
213
|
)
|
|
214
|
+
if managed := _managed_router_tokens(router_args or {}):
|
|
215
|
+
raise ValueError(
|
|
216
|
+
f"router_args cannot set PDEndpoint-managed flags: {managed}"
|
|
217
|
+
)
|
|
159
218
|
roles = ("prefill",) * ratio[0] + ("decode",) * ratio[1]
|
|
160
219
|
cluster = get_cluster_info()
|
|
161
220
|
hosts = cluster.container_ips
|
|
@@ -164,6 +223,13 @@ class PDEndpoint(Endpoint):
|
|
|
164
223
|
f"P/D ratio {ratio} requires a {len(roles)}-container Modal cluster"
|
|
165
224
|
)
|
|
166
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")
|
|
167
233
|
self.role = roles[cluster.rank]
|
|
168
234
|
super().__init__(_url(hosts[0], router_port))
|
|
169
235
|
self.engine = engine
|
|
@@ -199,16 +265,22 @@ class PDEndpoint(Endpoint):
|
|
|
199
265
|
health_interval=health_interval,
|
|
200
266
|
health_request_timeout=router_health_request_timeout,
|
|
201
267
|
health_failure_threshold=health_failure_threshold,
|
|
268
|
+
router_args=router_args or {},
|
|
202
269
|
)
|
|
203
270
|
if cluster.rank == 0
|
|
204
271
|
else None
|
|
205
272
|
)
|
|
273
|
+
# Fail before any weights load, not after a 20-minute engine boot.
|
|
274
|
+
if self.router is not None and not allow_custom_router:
|
|
275
|
+
verify_router()
|
|
206
276
|
self.drain_timeout = drain_timeout
|
|
207
277
|
self.health_interval = health_interval
|
|
208
278
|
self.health_failure_timeout = health_failure_timeout
|
|
209
279
|
self._stopped = threading.Event()
|
|
210
280
|
self._lock = threading.Lock()
|
|
211
281
|
self._started = False
|
|
282
|
+
self._serving = False
|
|
283
|
+
self._shutdown_listener: socket.socket | None = None
|
|
212
284
|
self._threads: list[threading.Thread] = []
|
|
213
285
|
|
|
214
286
|
def start(self, *, warmup: Callable[[], None] | None = None) -> None:
|
|
@@ -223,6 +295,13 @@ class PDEndpoint(Endpoint):
|
|
|
223
295
|
SGLANG_ENABLE_HEALTH_ENDPOINT_GENERATION="1",
|
|
224
296
|
)
|
|
225
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)
|
|
226
305
|
for endpoint in (self.engine, self.router):
|
|
227
306
|
if endpoint is not None:
|
|
228
307
|
endpoint.start()
|
|
@@ -234,6 +313,7 @@ class PDEndpoint(Endpoint):
|
|
|
234
313
|
for _, host in self.router.pd_config:
|
|
235
314
|
self._watch_health(_url(host, self.engine.worker_port) + "/health")
|
|
236
315
|
self._watch_health(self.router.base_url + "/health")
|
|
316
|
+
self._serving = True
|
|
237
317
|
except BaseException:
|
|
238
318
|
self.stop()
|
|
239
319
|
raise
|
|
@@ -295,7 +375,7 @@ class PDEndpoint(Endpoint):
|
|
|
295
375
|
if self._stopped.is_set():
|
|
296
376
|
return
|
|
297
377
|
self._stopped.set()
|
|
298
|
-
deadline = time.monotonic() +
|
|
378
|
+
deadline = time.monotonic() + _SHUTDOWN_TIMEOUT
|
|
299
379
|
try:
|
|
300
380
|
for endpoint in (self.router, self.engine):
|
|
301
381
|
if endpoint is None:
|
|
@@ -307,9 +387,58 @@ class PDEndpoint(Endpoint):
|
|
|
307
387
|
process.wait(timeout=max(0, deadline - time.monotonic()))
|
|
308
388
|
except subprocess.TimeoutExpired:
|
|
309
389
|
process.kill()
|
|
390
|
+
process.wait()
|
|
310
391
|
endpoint.stop()
|
|
311
392
|
finally:
|
|
312
393
|
for thread in self._threads:
|
|
313
|
-
thread.join(
|
|
314
|
-
|
|
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,
|
|
315
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])
|
|
@@ -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()
|
|
File without changes
|
|
File without changes
|
{autoinference_utils-0.2.8 → autoinference_utils-0.2.10}/src/autoinference_utils/__init__.py
RENAMED
|
File without changes
|