autoinference-utils 0.2.8__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.8 → autoinference_utils-0.2.9}/PKG-INFO +1 -1
- {autoinference_utils-0.2.8 → autoinference_utils-0.2.9}/pyproject.toml +1 -1
- {autoinference_utils-0.2.8 → autoinference_utils-0.2.9}/src/autoinference_utils/endpoint.py +94 -104
- {autoinference_utils-0.2.8 → autoinference_utils-0.2.9}/src/autoinference_utils/pd.py +68 -11
- autoinference_utils-0.2.9/src/autoinference_utils/router.py +103 -0
- {autoinference_utils-0.2.8 → autoinference_utils-0.2.9}/.gitignore +0 -0
- {autoinference_utils-0.2.8 → autoinference_utils-0.2.9}/README.md +0 -0
- {autoinference_utils-0.2.8 → autoinference_utils-0.2.9}/src/autoinference_utils/__init__.py +0 -0
|
@@ -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
|
|
|
@@ -19,7 +19,49 @@ from .endpoint import (
|
|
|
19
19
|
_local_base_url,
|
|
20
20
|
_merge_server_args,
|
|
21
21
|
_url,
|
|
22
|
+
server_arg_tokens,
|
|
22
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
|
+
)
|
|
23
65
|
|
|
24
66
|
|
|
25
67
|
def _configure_fabric() -> None:
|
|
@@ -61,6 +103,7 @@ class _PDRouter(RouterEndpoint):
|
|
|
61
103
|
health_interval,
|
|
62
104
|
health_request_timeout,
|
|
63
105
|
health_failure_threshold,
|
|
106
|
+
router_args,
|
|
64
107
|
**kwargs,
|
|
65
108
|
):
|
|
66
109
|
super().__init__(health_path="/readiness", host="::", **kwargs)
|
|
@@ -69,6 +112,7 @@ class _PDRouter(RouterEndpoint):
|
|
|
69
112
|
self.health_interval = health_interval
|
|
70
113
|
self.health_request_timeout = health_request_timeout
|
|
71
114
|
self.health_failure_threshold = health_failure_threshold
|
|
115
|
+
self.router_args = dict(router_args)
|
|
72
116
|
|
|
73
117
|
def _build_cmd(self) -> list[str]:
|
|
74
118
|
cmd = [
|
|
@@ -85,20 +129,21 @@ class _PDRouter(RouterEndpoint):
|
|
|
85
129
|
cmd.extend([f"--{role}", _url(host, self.worker_port)])
|
|
86
130
|
if role == "prefill":
|
|
87
131
|
cmd.append(str(self.prefill_bootstrap_port))
|
|
88
|
-
|
|
132
|
+
defaults = {
|
|
89
133
|
"--prefill-policy": "cache_aware",
|
|
90
134
|
"--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,
|
|
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",
|
|
99
143
|
"--log-level": "info",
|
|
100
|
-
}
|
|
101
|
-
|
|
144
|
+
}
|
|
145
|
+
for flag, value in _merge_server_args(defaults, self.router_args).items():
|
|
146
|
+
cmd.extend(server_arg_tokens(flag, value))
|
|
102
147
|
if self.api_key is not None:
|
|
103
148
|
cmd.extend(["--api-key", self.api_key])
|
|
104
149
|
return cmd
|
|
@@ -110,6 +155,8 @@ class PDEndpoint(Endpoint):
|
|
|
110
155
|
Call start/stop from Modal enter/exit hooks. Model-specific arguments and
|
|
111
156
|
warmup stay in the recipe. The image must support NIXL and graceful shutdown.
|
|
112
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.
|
|
113
160
|
"""
|
|
114
161
|
|
|
115
162
|
def __init__(
|
|
@@ -128,6 +175,8 @@ class PDEndpoint(Endpoint):
|
|
|
128
175
|
router_health_request_timeout: int = 5,
|
|
129
176
|
health_failure_threshold: int = 3,
|
|
130
177
|
health_failure_timeout: float = 30,
|
|
178
|
+
router_args: Mapping[str, str] | None = None,
|
|
179
|
+
allow_custom_router: bool = False,
|
|
131
180
|
):
|
|
132
181
|
from modal.experimental import get_cluster_info
|
|
133
182
|
|
|
@@ -156,6 +205,10 @@ class PDEndpoint(Endpoint):
|
|
|
156
205
|
raise ValueError(
|
|
157
206
|
"ratio must contain positive integer prefill and decode counts"
|
|
158
207
|
)
|
|
208
|
+
if managed := _managed_router_tokens(router_args or {}):
|
|
209
|
+
raise ValueError(
|
|
210
|
+
f"router_args cannot set PDEndpoint-managed flags: {managed}"
|
|
211
|
+
)
|
|
159
212
|
roles = ("prefill",) * ratio[0] + ("decode",) * ratio[1]
|
|
160
213
|
cluster = get_cluster_info()
|
|
161
214
|
hosts = cluster.container_ips
|
|
@@ -199,10 +252,14 @@ class PDEndpoint(Endpoint):
|
|
|
199
252
|
health_interval=health_interval,
|
|
200
253
|
health_request_timeout=router_health_request_timeout,
|
|
201
254
|
health_failure_threshold=health_failure_threshold,
|
|
255
|
+
router_args=router_args or {},
|
|
202
256
|
)
|
|
203
257
|
if cluster.rank == 0
|
|
204
258
|
else None
|
|
205
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()
|
|
206
263
|
self.drain_timeout = drain_timeout
|
|
207
264
|
self.health_interval = health_interval
|
|
208
265
|
self.health_failure_timeout = health_failure_timeout
|
|
@@ -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
|
|
File without changes
|