autoinference-utils 0.2.6__tar.gz → 0.2.7__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.7}/.gitignore +18 -1
- {autoinference_utils-0.2.6 → autoinference_utils-0.2.7}/PKG-INFO +7 -1
- autoinference_utils-0.2.7/README.md +15 -0
- {autoinference_utils-0.2.6 → autoinference_utils-0.2.7}/pyproject.toml +3 -3
- {autoinference_utils-0.2.6 → autoinference_utils-0.2.7/src}/autoinference_utils/endpoint.py +103 -110
- autoinference_utils-0.2.7/src/autoinference_utils/model_path.py +38 -0
- autoinference_utils-0.2.7/src/autoinference_utils/tests/test_model_path.py +84 -0
- autoinference_utils-0.2.6/README.md +0 -9
- {autoinference_utils-0.2.6 → autoinference_utils-0.2.7/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
|
Metadata-Version: 2.5
|
|
2
2
|
Name: autoinference-utils
|
|
3
|
-
Version: 0.2.
|
|
3
|
+
Version: 0.2.7
|
|
4
4
|
Summary: Shared endpoint abstractions for autoinference deployments
|
|
5
5
|
Requires-Python: >=3.10
|
|
6
6
|
Description-Content-Type: text/markdown
|
|
@@ -9,6 +9,12 @@ Description-Content-Type: text/markdown
|
|
|
9
9
|
|
|
10
10
|
Shared endpoint abstractions (`SGLangEndpoint`, `VLLMEndpoint`) for autoinference deployments.
|
|
11
11
|
|
|
12
|
+
For SGLang models whose custom Transformers code cannot load from symlinked
|
|
13
|
+
Hugging Face snapshots, set `AUTOINFERENCE_MATERIALIZE_MODEL_PATH=1`. The endpoint
|
|
14
|
+
copies local base/draft model code and metadata into unique temporary directories
|
|
15
|
+
and symlinks top-level safetensors weights. Repository IDs pass through unchanged.
|
|
16
|
+
Leave the setting unset to preserve the original paths.
|
|
17
|
+
|
|
12
18
|
## Publishing
|
|
13
19
|
|
|
14
20
|
Bump `version` in `pyproject.toml` and merge to `main` — CI publishes automatically via trusted publishing.
|
|
@@ -0,0 +1,15 @@
|
|
|
1
|
+
# autoinference-utils
|
|
2
|
+
|
|
3
|
+
Shared endpoint abstractions (`SGLangEndpoint`, `VLLMEndpoint`) for autoinference deployments.
|
|
4
|
+
|
|
5
|
+
For SGLang models whose custom Transformers code cannot load from symlinked
|
|
6
|
+
Hugging Face snapshots, set `AUTOINFERENCE_MATERIALIZE_MODEL_PATH=1`. The endpoint
|
|
7
|
+
copies local base/draft model code and metadata into unique temporary directories
|
|
8
|
+
and symlinks top-level safetensors weights. Repository IDs pass through unchanged.
|
|
9
|
+
Leave the setting unset to preserve the original paths.
|
|
10
|
+
|
|
11
|
+
## Publishing
|
|
12
|
+
|
|
13
|
+
Bump `version` in `pyproject.toml` and merge to `main` — CI publishes automatically via trusted publishing.
|
|
14
|
+
|
|
15
|
+
To publish manually: `cd autoinference_utils && uv build && uv publish`
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
[project]
|
|
2
2
|
name = "autoinference-utils"
|
|
3
|
-
version = "0.2.
|
|
3
|
+
version = "0.2.7"
|
|
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"]
|
|
@@ -23,11 +23,13 @@ import urllib.request
|
|
|
23
23
|
from abc import ABC
|
|
24
24
|
from typing import Any, Callable, Literal, Mapping, Optional, Sequence
|
|
25
25
|
|
|
26
|
+
from autoinference_utils.model_path import materialize_model_path
|
|
27
|
+
|
|
26
28
|
BENCH_MODE_ENV = "ENDPOINT_BENCH_MODE"
|
|
27
29
|
DUMMY_WEIGHTS_ENV = "ENDPOINT_DUMMY"
|
|
28
30
|
BENCH_MODE_PORT_OFFSET = 10000
|
|
29
31
|
|
|
30
|
-
MODAL_FLASH_REQUEST_UUID_HEADER= "modal-flash-request-uuid"
|
|
32
|
+
MODAL_FLASH_REQUEST_UUID_HEADER = "modal-flash-request-uuid"
|
|
31
33
|
MODAL_SESSION_ID_HEADER = "modal-session-id"
|
|
32
34
|
|
|
33
35
|
ENDPOINTS_REQUIRING_OAI_STREAM_COMPAT: tuple[str, ...] = ("TRTLLMEndpoint",)
|
|
@@ -55,9 +57,7 @@ def _flag_name(key: str) -> str:
|
|
|
55
57
|
return key.partition("=")[0]
|
|
56
58
|
|
|
57
59
|
|
|
58
|
-
def get_server_arg(
|
|
59
|
-
server_args: Mapping[str, str], flag: str
|
|
60
|
-
) -> tuple[str, str] | None:
|
|
60
|
+
def get_server_arg(server_args: Mapping[str, str], flag: str) -> tuple[str, str] | None:
|
|
61
61
|
"""Return the key and value setting a flag under either spelling."""
|
|
62
62
|
for key, value in server_args.items():
|
|
63
63
|
if _flag_name(key) == flag:
|
|
@@ -153,11 +153,15 @@ class SGLangEndpoint(Endpoint):
|
|
|
153
153
|
super().__init__(base_url=f"http://localhost:{sglang_port}")
|
|
154
154
|
|
|
155
155
|
self.worker_port = sglang_port
|
|
156
|
-
self.model_path = model_path
|
|
156
|
+
self.model_path = materialize_model_path(model_path)
|
|
157
157
|
self.tp = tp
|
|
158
158
|
self.ep = ep
|
|
159
159
|
self.dp = dp
|
|
160
|
-
self.speculative_model_path =
|
|
160
|
+
self.speculative_model_path = (
|
|
161
|
+
materialize_model_path(speculative_model_path)
|
|
162
|
+
if speculative_model_path is not None
|
|
163
|
+
else None
|
|
164
|
+
)
|
|
161
165
|
self.load_format = load_format
|
|
162
166
|
self.nnodes = nnodes
|
|
163
167
|
self.node_rank = node_rank
|
|
@@ -166,9 +170,7 @@ class SGLangEndpoint(Endpoint):
|
|
|
166
170
|
self.disaggregation_mode = disaggregation_mode
|
|
167
171
|
self.prefill_bootstrap_port = prefill_bootstrap_port
|
|
168
172
|
self.launcher_module = launcher_module
|
|
169
|
-
self.extra_server_args = (
|
|
170
|
-
dict(extra_server_args) if extra_server_args else {}
|
|
171
|
-
)
|
|
173
|
+
self.extra_server_args = dict(extra_server_args) if extra_server_args else {}
|
|
172
174
|
self.health_timeout = health_timeout
|
|
173
175
|
self.health_poll_interval = health_poll_interval
|
|
174
176
|
self.health_request_timeout = health_request_timeout
|
|
@@ -183,17 +185,16 @@ class SGLangEndpoint(Endpoint):
|
|
|
183
185
|
self._bench_server: Optional[http.server.ThreadingHTTPServer] = None
|
|
184
186
|
|
|
185
187
|
if self.disaggregation_mode not in (None, "prefill", "decode"):
|
|
186
|
-
raise ValueError(
|
|
187
|
-
"disaggregation_mode must be None, 'prefill', or 'decode'"
|
|
188
|
-
)
|
|
188
|
+
raise ValueError("disaggregation_mode must be None, 'prefill', or 'decode'")
|
|
189
189
|
|
|
190
190
|
if os.environ.get(DUMMY_WEIGHTS_ENV) == "1":
|
|
191
|
-
if (
|
|
192
|
-
self.
|
|
193
|
-
and not has_server_arg(self.extra_server_args, "--load-format")
|
|
191
|
+
if self.load_format is None and not has_server_arg(
|
|
192
|
+
self.extra_server_args, "--load-format"
|
|
194
193
|
):
|
|
195
194
|
self.load_format = "dummy"
|
|
196
|
-
if not has_server_arg(
|
|
195
|
+
if not has_server_arg(
|
|
196
|
+
self.extra_server_args, "--model-loader-extra-config"
|
|
197
|
+
):
|
|
197
198
|
self.extra_server_args["--model-loader-extra-config"] = "{}"
|
|
198
199
|
|
|
199
200
|
# Request logging enabled, log only metadata by default
|
|
@@ -204,7 +205,9 @@ class SGLangEndpoint(Endpoint):
|
|
|
204
205
|
# Set request logging level
|
|
205
206
|
level = get_server_arg(self.extra_server_args, "--log-requests-level")
|
|
206
207
|
if level is None:
|
|
207
|
-
self.extra_server_args["--log-requests-level"] = str(
|
|
208
|
+
self.extra_server_args["--log-requests-level"] = str(
|
|
209
|
+
self.log_requests_level
|
|
210
|
+
)
|
|
208
211
|
else:
|
|
209
212
|
key, value = level
|
|
210
213
|
print(
|
|
@@ -226,19 +229,21 @@ class SGLangEndpoint(Endpoint):
|
|
|
226
229
|
self.log_request_headers = ",".join(sglang_log_request_headers)
|
|
227
230
|
os.environ["SGLANG_LOG_REQUEST_HEADERS"] = self.log_request_headers
|
|
228
231
|
|
|
229
|
-
|
|
230
232
|
def _build_cmd(self) -> list[str]:
|
|
231
233
|
cmd = [
|
|
232
|
-
"python",
|
|
233
|
-
"
|
|
234
|
-
|
|
235
|
-
"--
|
|
234
|
+
"python",
|
|
235
|
+
"-m",
|
|
236
|
+
self.launcher_module,
|
|
237
|
+
"--host",
|
|
238
|
+
"0.0.0.0",
|
|
239
|
+
"--port",
|
|
240
|
+
str(self.worker_port),
|
|
241
|
+
"--model-path",
|
|
242
|
+
self.model_path,
|
|
236
243
|
]
|
|
237
244
|
|
|
238
245
|
if self.speculative_model_path is not None:
|
|
239
|
-
cmd.extend(
|
|
240
|
-
["--speculative-draft-model-path", self.speculative_model_path]
|
|
241
|
-
)
|
|
246
|
+
cmd.extend(["--speculative-draft-model-path", self.speculative_model_path])
|
|
242
247
|
if self.load_format is not None:
|
|
243
248
|
cmd.extend(["--load-format", self.load_format])
|
|
244
249
|
if self.tp is not None:
|
|
@@ -263,14 +268,18 @@ class SGLangEndpoint(Endpoint):
|
|
|
263
268
|
raise ValueError("dist_init_host is required when nnodes > 1")
|
|
264
269
|
cmd.extend(
|
|
265
270
|
[
|
|
266
|
-
"--nnodes",
|
|
267
|
-
|
|
271
|
+
"--nnodes",
|
|
272
|
+
str(self.nnodes),
|
|
273
|
+
"--node-rank",
|
|
274
|
+
str(self.node_rank),
|
|
268
275
|
"--dist-init-addr",
|
|
269
276
|
f"{self.dist_init_host}:{self.dist_init_port}",
|
|
270
277
|
]
|
|
271
278
|
)
|
|
272
279
|
|
|
273
|
-
merged = _merge_server_args(
|
|
280
|
+
merged = _merge_server_args(
|
|
281
|
+
self.DEFAULT_OPERATIONAL_ARGS, self.extra_server_args
|
|
282
|
+
)
|
|
274
283
|
for key, value in merged.items():
|
|
275
284
|
cmd.extend(server_arg_tokens(key, value))
|
|
276
285
|
|
|
@@ -333,9 +342,7 @@ class VLLMEndpoint(Endpoint):
|
|
|
333
342
|
super().__init__(base_url=f"http://localhost:{vllm_port}")
|
|
334
343
|
self.model = model
|
|
335
344
|
self.worker_port = vllm_port
|
|
336
|
-
self.extra_server_args = (
|
|
337
|
-
dict(extra_server_args) if extra_server_args else {}
|
|
338
|
-
)
|
|
345
|
+
self.extra_server_args = dict(extra_server_args) if extra_server_args else {}
|
|
339
346
|
self.health_timeout = health_timeout
|
|
340
347
|
self.health_poll_interval = health_poll_interval
|
|
341
348
|
self.health_request_timeout = health_request_timeout
|
|
@@ -348,10 +355,15 @@ class VLLMEndpoint(Endpoint):
|
|
|
348
355
|
|
|
349
356
|
def _build_cmd(self) -> list[str]:
|
|
350
357
|
cmd = [
|
|
351
|
-
"python",
|
|
352
|
-
"
|
|
353
|
-
"
|
|
354
|
-
"--
|
|
358
|
+
"python",
|
|
359
|
+
"-m",
|
|
360
|
+
"vllm.entrypoints.openai.api_server",
|
|
361
|
+
"--host",
|
|
362
|
+
"0.0.0.0",
|
|
363
|
+
"--port",
|
|
364
|
+
str(self.worker_port),
|
|
365
|
+
"--model",
|
|
366
|
+
self.model,
|
|
355
367
|
]
|
|
356
368
|
for key, value in self.extra_server_args.items():
|
|
357
369
|
cmd.extend(server_arg_tokens(key, value))
|
|
@@ -405,9 +417,7 @@ class TRTLLMEndpoint(Endpoint):
|
|
|
405
417
|
super().__init__(base_url=f"http://localhost:{worker_port}")
|
|
406
418
|
self.model = model
|
|
407
419
|
self.worker_port = worker_port
|
|
408
|
-
self.extra_server_args = (
|
|
409
|
-
dict(extra_server_args) if extra_server_args else {}
|
|
410
|
-
)
|
|
420
|
+
self.extra_server_args = dict(extra_server_args) if extra_server_args else {}
|
|
411
421
|
self.health_timeout = health_timeout
|
|
412
422
|
self.health_poll_interval = health_poll_interval
|
|
413
423
|
self.health_request_timeout = health_request_timeout
|
|
@@ -417,8 +427,10 @@ class TRTLLMEndpoint(Endpoint):
|
|
|
417
427
|
cmd = [
|
|
418
428
|
"trtllm-serve",
|
|
419
429
|
self.model,
|
|
420
|
-
"--host",
|
|
421
|
-
"
|
|
430
|
+
"--host",
|
|
431
|
+
"0.0.0.0",
|
|
432
|
+
"--port",
|
|
433
|
+
str(self.worker_port),
|
|
422
434
|
]
|
|
423
435
|
for key, value in self.extra_server_args.items():
|
|
424
436
|
cmd.extend(server_arg_tokens(key, value))
|
|
@@ -483,18 +495,30 @@ class RouterEndpoint(Endpoint):
|
|
|
483
495
|
|
|
484
496
|
def _build_cmd(self) -> list[str]:
|
|
485
497
|
cmd = [
|
|
486
|
-
"python",
|
|
487
|
-
"
|
|
488
|
-
"
|
|
489
|
-
"--
|
|
490
|
-
"
|
|
491
|
-
"--
|
|
492
|
-
|
|
493
|
-
"--
|
|
494
|
-
"
|
|
495
|
-
"--
|
|
498
|
+
"python",
|
|
499
|
+
"-m",
|
|
500
|
+
"sglang_router.launch_router",
|
|
501
|
+
"--host",
|
|
502
|
+
"0.0.0.0",
|
|
503
|
+
"--port",
|
|
504
|
+
str(self.router_port),
|
|
505
|
+
"--prefill-policy",
|
|
506
|
+
"cache_aware",
|
|
507
|
+
"--decode-policy",
|
|
508
|
+
"round_robin",
|
|
509
|
+
"--max-concurrent-requests",
|
|
510
|
+
"128",
|
|
511
|
+
"--rate-limit-tokens-per-second",
|
|
512
|
+
"0",
|
|
513
|
+
"--queue-size",
|
|
514
|
+
"0",
|
|
515
|
+
"--health-check-timeout-secs",
|
|
516
|
+
"600",
|
|
517
|
+
"--log-level",
|
|
518
|
+
"info",
|
|
496
519
|
"--disable-circuit-breaker",
|
|
497
|
-
"--request-timeout-secs",
|
|
520
|
+
"--request-timeout-secs",
|
|
521
|
+
"3600",
|
|
498
522
|
]
|
|
499
523
|
|
|
500
524
|
if self.api_key is not None:
|
|
@@ -506,9 +530,7 @@ class RouterEndpoint(Endpoint):
|
|
|
506
530
|
for role, node_ip in self.pd_config:
|
|
507
531
|
node_url = f"http://{node_ip}:{self.worker_port}"
|
|
508
532
|
if role == "prefill":
|
|
509
|
-
cmd.extend(
|
|
510
|
-
["--prefill", node_url, str(self.prefill_bootstrap_port)]
|
|
511
|
-
)
|
|
533
|
+
cmd.extend(["--prefill", node_url, str(self.prefill_bootstrap_port)])
|
|
512
534
|
elif role == "decode":
|
|
513
535
|
cmd.extend(["--decode", node_url])
|
|
514
536
|
elif role == "worker":
|
|
@@ -552,6 +574,7 @@ class RouterEndpoint(Endpoint):
|
|
|
552
574
|
# and this proxy listens on worker_port. POST /bench shells out to a benchmark
|
|
553
575
|
# task via run_bench(); every other path is forwarded to the upstream server.
|
|
554
576
|
|
|
577
|
+
|
|
555
578
|
def start_bench_proxy(
|
|
556
579
|
*,
|
|
557
580
|
listen_port: int,
|
|
@@ -562,17 +585,12 @@ def start_bench_proxy(
|
|
|
562
585
|
(_BenchProxyHandler,),
|
|
563
586
|
{"upstream_port": upstream_port},
|
|
564
587
|
)
|
|
565
|
-
server = http.server.ThreadingHTTPServer(
|
|
566
|
-
("0.0.0.0", listen_port), handler_cls
|
|
567
|
-
)
|
|
588
|
+
server = http.server.ThreadingHTTPServer(("0.0.0.0", listen_port), handler_cls)
|
|
568
589
|
thread = threading.Thread(
|
|
569
590
|
target=server.serve_forever, daemon=True, name="bench-proxy"
|
|
570
591
|
)
|
|
571
592
|
thread.start()
|
|
572
|
-
print(
|
|
573
|
-
f"[bench-proxy] listening on :{listen_port}, "
|
|
574
|
-
f"forwarding to :{upstream_port}"
|
|
575
|
-
)
|
|
593
|
+
print(f"[bench-proxy] listening on :{listen_port}, forwarding to :{upstream_port}")
|
|
576
594
|
return server
|
|
577
595
|
|
|
578
596
|
|
|
@@ -601,26 +619,17 @@ class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
|
|
|
601
619
|
try:
|
|
602
620
|
payload = self._read_json_body()
|
|
603
621
|
except (ValueError, json.JSONDecodeError) as exc:
|
|
604
|
-
return self._send_json(
|
|
605
|
-
400, {"ok": False, "error": f"bad body: {exc}"}
|
|
606
|
-
)
|
|
622
|
+
return self._send_json(400, {"ok": False, "error": f"bad body: {exc}"})
|
|
607
623
|
|
|
608
624
|
benchmark = payload.get("benchmark") or ""
|
|
609
625
|
args = payload.get("args") or []
|
|
610
|
-
target = (
|
|
611
|
-
payload.get("target")
|
|
612
|
-
or f"http://localhost:{self.upstream_port}"
|
|
613
|
-
)
|
|
626
|
+
target = payload.get("target") or f"http://localhost:{self.upstream_port}"
|
|
614
627
|
output_dir = payload.get("output_dir") or "/tmp/bench-output"
|
|
615
628
|
|
|
616
629
|
if not benchmark:
|
|
617
|
-
return self._send_json(
|
|
618
|
-
400, {"ok": False, "error": "benchmark is required"}
|
|
619
|
-
)
|
|
630
|
+
return self._send_json(400, {"ok": False, "error": "benchmark is required"})
|
|
620
631
|
if not isinstance(args, list):
|
|
621
|
-
return self._send_json(
|
|
622
|
-
400, {"ok": False, "error": "args must be a list"}
|
|
623
|
-
)
|
|
632
|
+
return self._send_json(400, {"ok": False, "error": "args must be a list"})
|
|
624
633
|
|
|
625
634
|
rc, body = run_bench(benchmark, list(args), target, output_dir)
|
|
626
635
|
status = 200 if rc == 0 else 500
|
|
@@ -671,9 +680,7 @@ class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
|
|
|
671
680
|
return json.loads(raw)
|
|
672
681
|
|
|
673
682
|
def _send_json(self, status: int, obj: Mapping[str, Any]):
|
|
674
|
-
self._send_raw(
|
|
675
|
-
status, "application/json", json.dumps(obj).encode()
|
|
676
|
-
)
|
|
683
|
+
self._send_raw(status, "application/json", json.dumps(obj).encode())
|
|
677
684
|
|
|
678
685
|
def _send_raw(self, status: int, content_type: str, body: bytes):
|
|
679
686
|
self.send_response(status)
|
|
@@ -698,18 +705,21 @@ def run_bench(
|
|
|
698
705
|
return 500, json.dumps(
|
|
699
706
|
{
|
|
700
707
|
"ok": False,
|
|
701
|
-
"error": (
|
|
702
|
-
"autoinference.benchmarks not importable in container"
|
|
703
|
-
),
|
|
708
|
+
"error": ("autoinference.benchmarks not importable in container"),
|
|
704
709
|
"volume_path": output_dir,
|
|
705
710
|
}
|
|
706
711
|
).encode()
|
|
707
712
|
|
|
708
713
|
cmd = [
|
|
709
|
-
sys.executable,
|
|
710
|
-
"
|
|
711
|
-
"
|
|
712
|
-
|
|
714
|
+
sys.executable,
|
|
715
|
+
"-m",
|
|
716
|
+
"invoke",
|
|
717
|
+
"--search-root",
|
|
718
|
+
search_root,
|
|
719
|
+
"-c",
|
|
720
|
+
"tasks",
|
|
721
|
+
benchmark,
|
|
722
|
+
*[str(a) for a in args],
|
|
713
723
|
]
|
|
714
724
|
env = {
|
|
715
725
|
**os.environ,
|
|
@@ -759,9 +769,7 @@ def wait_ready(
|
|
|
759
769
|
|
|
760
770
|
while time.time() < deadline:
|
|
761
771
|
try:
|
|
762
|
-
error = _health_check(
|
|
763
|
-
url, request_timeout=request_timeout, process=process
|
|
764
|
-
)
|
|
772
|
+
error = _health_check(url, request_timeout=request_timeout, process=process)
|
|
765
773
|
except subprocess.CalledProcessError as exc:
|
|
766
774
|
print(
|
|
767
775
|
f"[endpoint] !!! server process exited with code "
|
|
@@ -775,13 +783,11 @@ def wait_ready(
|
|
|
775
783
|
time.sleep(poll_interval)
|
|
776
784
|
|
|
777
785
|
print(
|
|
778
|
-
f"[endpoint] !!! health check timed out after {timeout}s; "
|
|
779
|
-
f"last={last_error}",
|
|
786
|
+
f"[endpoint] !!! health check timed out after {timeout}s; last={last_error}",
|
|
780
787
|
flush=True,
|
|
781
788
|
)
|
|
782
789
|
raise TimeoutError(
|
|
783
|
-
f"Health check timed out after {timeout}s for {url}. "
|
|
784
|
-
f"Last error: {last_error}"
|
|
790
|
+
f"Health check timed out after {timeout}s for {url}. Last error: {last_error}"
|
|
785
791
|
)
|
|
786
792
|
|
|
787
793
|
|
|
@@ -809,9 +815,7 @@ def warmup_chat_completions(
|
|
|
809
815
|
request_timeout=request_timeout,
|
|
810
816
|
max_attempts=max_attempts_per_request,
|
|
811
817
|
retry_delay=retry_delay,
|
|
812
|
-
description=(
|
|
813
|
-
f"warmup request {request_idx + 1}/{successful_requests}"
|
|
814
|
-
),
|
|
818
|
+
description=(f"warmup request {request_idx + 1}/{successful_requests}"),
|
|
815
819
|
)
|
|
816
820
|
|
|
817
821
|
|
|
@@ -833,9 +837,7 @@ def validate_embeddings_endpoint(
|
|
|
833
837
|
1 if all(isinstance(item, int) for item in inputs) else len(inputs)
|
|
834
838
|
)
|
|
835
839
|
else:
|
|
836
|
-
raise ValueError(
|
|
837
|
-
"embedding validation payload must contain non-empty input"
|
|
838
|
-
)
|
|
840
|
+
raise ValueError("embedding validation payload must contain non-empty input")
|
|
839
841
|
|
|
840
842
|
request_headers = {"Content-Type": "application/json"}
|
|
841
843
|
if headers:
|
|
@@ -898,10 +900,7 @@ def start_heartbeat_thread(
|
|
|
898
900
|
try:
|
|
899
901
|
error = health_check_fn()
|
|
900
902
|
except subprocess.CalledProcessError as exc:
|
|
901
|
-
print(
|
|
902
|
-
f"[heartbeat] server process exited with code "
|
|
903
|
-
f"{exc.returncode}"
|
|
904
|
-
)
|
|
903
|
+
print(f"[heartbeat] server process exited with code {exc.returncode}")
|
|
905
904
|
on_failure()
|
|
906
905
|
return
|
|
907
906
|
if error is None:
|
|
@@ -913,10 +912,7 @@ def start_heartbeat_thread(
|
|
|
913
912
|
f"({consecutive_failures}/{max_consecutive_failures})"
|
|
914
913
|
)
|
|
915
914
|
if consecutive_failures >= max_consecutive_failures:
|
|
916
|
-
print(
|
|
917
|
-
"[heartbeat] sustained health-check failure, "
|
|
918
|
-
"invoking on_failure"
|
|
919
|
-
)
|
|
915
|
+
print("[heartbeat] sustained health-check failure, invoking on_failure")
|
|
920
916
|
on_failure()
|
|
921
917
|
return
|
|
922
918
|
|
|
@@ -1001,8 +997,7 @@ def _wait_ready_url(
|
|
|
1001
997
|
last_error = error
|
|
1002
998
|
time.sleep(poll_interval)
|
|
1003
999
|
raise TimeoutError(
|
|
1004
|
-
f"Timed out after {timeout}s waiting for {url}. "
|
|
1005
|
-
f"Last error: {last_error}"
|
|
1000
|
+
f"Timed out after {timeout}s waiting for {url}. Last error: {last_error}"
|
|
1006
1001
|
)
|
|
1007
1002
|
|
|
1008
1003
|
|
|
@@ -1020,9 +1015,7 @@ def _post_json(
|
|
|
1020
1015
|
timeout: float | None = None,
|
|
1021
1016
|
) -> bytes:
|
|
1022
1017
|
body = json.dumps(payload).encode("utf-8")
|
|
1023
|
-
req = urllib.request.Request(
|
|
1024
|
-
url, data=body, headers=dict(headers), method="POST"
|
|
1025
|
-
)
|
|
1018
|
+
req = urllib.request.Request(url, data=body, headers=dict(headers), method="POST")
|
|
1026
1019
|
with urllib.request.urlopen(req, timeout=timeout) as resp:
|
|
1027
1020
|
return resp.read()
|
|
1028
1021
|
|
|
@@ -0,0 +1,38 @@
|
|
|
1
|
+
"""Prepare local model snapshots for Transformers' remote-code loader."""
|
|
2
|
+
|
|
3
|
+
import os
|
|
4
|
+
import shutil
|
|
5
|
+
import tempfile
|
|
6
|
+
from pathlib import Path
|
|
7
|
+
|
|
8
|
+
MATERIALIZE_MODEL_PATH_ENV = "AUTOINFERENCE_MATERIALIZE_MODEL_PATH"
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
def materialize_model_path(path: str) -> str:
|
|
12
|
+
"""Copy snapshot metadata/code into a real directory; retain weight links.
|
|
13
|
+
|
|
14
|
+
Transformers can resolve Python-module symlinks into the HF blob directory,
|
|
15
|
+
losing the relative imports next to the original snapshot entry. This opt-in
|
|
16
|
+
keeps small files together without copying large safetensors shards. The
|
|
17
|
+
unique directory lasts for the container's lifetime.
|
|
18
|
+
"""
|
|
19
|
+
source = Path(path)
|
|
20
|
+
if os.environ.get(MATERIALIZE_MODEL_PATH_ENV) != "1" or not source.is_dir():
|
|
21
|
+
return path
|
|
22
|
+
|
|
23
|
+
source = source.resolve()
|
|
24
|
+
destination = Path(tempfile.mkdtemp(prefix="autoinference-model-"))
|
|
25
|
+
try:
|
|
26
|
+
for entry in source.iterdir():
|
|
27
|
+
target = destination / entry.name
|
|
28
|
+
if entry.is_file() and entry.suffix == ".safetensors":
|
|
29
|
+
target.symlink_to(entry.resolve(strict=True))
|
|
30
|
+
elif entry.is_dir():
|
|
31
|
+
shutil.copytree(entry, target, symlinks=False)
|
|
32
|
+
else:
|
|
33
|
+
shutil.copy2(entry, target, follow_symlinks=True)
|
|
34
|
+
except BaseException:
|
|
35
|
+
shutil.rmtree(destination)
|
|
36
|
+
raise
|
|
37
|
+
print(f"Materialized model snapshot {source} -> {destination}")
|
|
38
|
+
return str(destination)
|
|
@@ -0,0 +1,84 @@
|
|
|
1
|
+
import shutil
|
|
2
|
+
from pathlib import Path
|
|
3
|
+
|
|
4
|
+
import pytest
|
|
5
|
+
|
|
6
|
+
from autoinference_utils.endpoint import SGLangEndpoint
|
|
7
|
+
from autoinference_utils.model_path import (
|
|
8
|
+
MATERIALIZE_MODEL_PATH_ENV,
|
|
9
|
+
materialize_model_path,
|
|
10
|
+
)
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def snapshot(tmp_path):
|
|
14
|
+
blobs = tmp_path / "blobs"
|
|
15
|
+
source = tmp_path / "snapshots" / "revision"
|
|
16
|
+
blobs.mkdir()
|
|
17
|
+
source.mkdir(parents=True)
|
|
18
|
+
for name, content in {
|
|
19
|
+
"config.json": '{"model_type": "kimi"}',
|
|
20
|
+
"configuration_kimi.py": "from .helper import VALUE\n",
|
|
21
|
+
"helper.py": "VALUE = 42\n",
|
|
22
|
+
"model.safetensors": "weight placeholder",
|
|
23
|
+
}.items():
|
|
24
|
+
blob = blobs / ("hash-" + name)
|
|
25
|
+
blob.write_text(content)
|
|
26
|
+
(source / name).symlink_to(blob)
|
|
27
|
+
return source
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
def test_opt_in_and_repository_ids(tmp_path, monkeypatch):
|
|
31
|
+
source = snapshot(tmp_path)
|
|
32
|
+
monkeypatch.delenv(MATERIALIZE_MODEL_PATH_ENV, raising=False)
|
|
33
|
+
assert materialize_model_path(str(source)) == str(source)
|
|
34
|
+
monkeypatch.setenv(MATERIALIZE_MODEL_PATH_ENV, "1")
|
|
35
|
+
assert materialize_model_path("moonshotai/Kimi-K3") == "moonshotai/Kimi-K3"
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
def test_local_code_is_materialized_and_weights_stay_linked(tmp_path, monkeypatch):
|
|
39
|
+
source = snapshot(tmp_path)
|
|
40
|
+
monkeypatch.setenv(MATERIALIZE_MODEL_PATH_ENV, "1")
|
|
41
|
+
monkeypatch.setattr("tempfile.tempdir", str(tmp_path))
|
|
42
|
+
result = Path(materialize_model_path(str(source)))
|
|
43
|
+
assert result != source
|
|
44
|
+
for name in ("config.json", "configuration_kimi.py", "helper.py"):
|
|
45
|
+
assert not (result / name).is_symlink()
|
|
46
|
+
assert (result / name).resolve().parent == result
|
|
47
|
+
assert (result / name).read_bytes() == (source / name).read_bytes()
|
|
48
|
+
assert (result / "model.safetensors").is_symlink()
|
|
49
|
+
assert (result / "model.safetensors").resolve() == (
|
|
50
|
+
source / "model.safetensors"
|
|
51
|
+
).resolve()
|
|
52
|
+
assert (source / "configuration_kimi.py").is_symlink()
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
def test_endpoint_uses_materialized_base_and_draft(tmp_path, monkeypatch):
|
|
56
|
+
source = snapshot(tmp_path)
|
|
57
|
+
monkeypatch.setenv(MATERIALIZE_MODEL_PATH_ENV, "1")
|
|
58
|
+
monkeypatch.setattr("tempfile.tempdir", str(tmp_path))
|
|
59
|
+
endpoint = SGLangEndpoint(
|
|
60
|
+
model_path=str(source), speculative_model_path=str(source)
|
|
61
|
+
)
|
|
62
|
+
assert endpoint.model_path != str(source)
|
|
63
|
+
assert endpoint.speculative_model_path != str(source)
|
|
64
|
+
assert endpoint.model_path != endpoint.speculative_model_path
|
|
65
|
+
cmd = endpoint._build_cmd()
|
|
66
|
+
assert cmd[cmd.index("--model-path") + 1] == endpoint.model_path
|
|
67
|
+
assert (
|
|
68
|
+
cmd[cmd.index("--speculative-draft-model-path") + 1]
|
|
69
|
+
== endpoint.speculative_model_path
|
|
70
|
+
)
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
def test_failed_copy_removes_partial_directory(tmp_path, monkeypatch):
|
|
74
|
+
source = snapshot(tmp_path)
|
|
75
|
+
monkeypatch.setenv(MATERIALIZE_MODEL_PATH_ENV, "1")
|
|
76
|
+
monkeypatch.setattr("tempfile.tempdir", str(tmp_path))
|
|
77
|
+
|
|
78
|
+
def fail(*args, **kwargs):
|
|
79
|
+
raise OSError("copy failed")
|
|
80
|
+
|
|
81
|
+
monkeypatch.setattr(shutil, "copy2", fail)
|
|
82
|
+
with pytest.raises(OSError, match="copy failed"):
|
|
83
|
+
materialize_model_path(str(source))
|
|
84
|
+
assert not list(tmp_path.glob("autoinference-model-*"))
|
|
@@ -1,9 +0,0 @@
|
|
|
1
|
-
# autoinference-utils
|
|
2
|
-
|
|
3
|
-
Shared endpoint abstractions (`SGLangEndpoint`, `VLLMEndpoint`) for autoinference deployments.
|
|
4
|
-
|
|
5
|
-
## Publishing
|
|
6
|
-
|
|
7
|
-
Bump `version` in `pyproject.toml` and merge to `main` — CI publishes automatically via trusted publishing.
|
|
8
|
-
|
|
9
|
-
To publish manually: `cd autoinference_utils && uv build && uv publish`
|
|
File without changes
|