autoinference-utils 0.2.5__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.5 → autoinference_utils-0.2.7}/.gitignore +21 -6
- {autoinference_utils-0.2.5 → autoinference_utils-0.2.7}/PKG-INFO +8 -2
- autoinference_utils-0.2.7/README.md +15 -0
- {autoinference_utils-0.2.5 → autoinference_utils-0.2.7}/pyproject.toml +3 -3
- {autoinference_utils-0.2.5 → autoinference_utils-0.2.7/src}/autoinference_utils/endpoint.py +163 -129
- 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.5/README.md +0 -9
- {autoinference_utils-0.2.5 → autoinference_utils-0.2.7/src}/autoinference_utils/__init__.py +0 -0
|
@@ -12,6 +12,7 @@ node_modules/
|
|
|
12
12
|
build/
|
|
13
13
|
develop-eggs/
|
|
14
14
|
dist/
|
|
15
|
+
!.github/actions/paths-filter/dist/
|
|
15
16
|
downloads/
|
|
16
17
|
eggs/
|
|
17
18
|
.eggs/
|
|
@@ -149,8 +150,10 @@ activemq-data/
|
|
|
149
150
|
|
|
150
151
|
# Environments
|
|
151
152
|
.env
|
|
153
|
+
.env.dev
|
|
152
154
|
.envrc
|
|
153
155
|
.venv
|
|
156
|
+
.autoinference/
|
|
154
157
|
env/
|
|
155
158
|
venv/
|
|
156
159
|
ENV/
|
|
@@ -239,10 +242,22 @@ Network Trash Folder
|
|
|
239
242
|
Temporary Items
|
|
240
243
|
.apdisk
|
|
241
244
|
|
|
245
|
+
# Ephemeral files from benchmark runs
|
|
242
246
|
benchmark_results/
|
|
243
|
-
|
|
244
|
-
|
|
245
|
-
|
|
246
|
-
|
|
247
|
-
|
|
248
|
-
|
|
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
|
-
Metadata-Version: 2.
|
|
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,19 +23,62 @@ 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
|
-
# Endpoint class names whose servers emit OAI streaming chunks that
|
|
34
|
-
# sglang.bench_serving can't parse. deploy_and_bench auto-injects
|
|
35
|
-
# --openai-stream-compat for deployments importing any of these.
|
|
36
35
|
ENDPOINTS_REQUIRING_OAI_STREAM_COMPAT: tuple[str, ...] = ("TRTLLMEndpoint",)
|
|
37
36
|
|
|
38
37
|
|
|
38
|
+
def server_arg_tokens(flag: str, value: str) -> list[str]:
|
|
39
|
+
"""Render one server arg as argv tokens.
|
|
40
|
+
|
|
41
|
+
A value splits on whitespace, so a repeatable flag takes its values as one
|
|
42
|
+
entry. A lone value starting with `-` keeps the attached `--flag=value`
|
|
43
|
+
spelling, the only one an argparse-style parser does not read as a flag.
|
|
44
|
+
"""
|
|
45
|
+
if value and "=" in flag:
|
|
46
|
+
raise ValueError(f"server arg {flag!r} already carries a value: {value!r}")
|
|
47
|
+
if value == "":
|
|
48
|
+
return [flag]
|
|
49
|
+
tokens = value.split()
|
|
50
|
+
if len(tokens) == 1 and tokens[0].startswith("-"):
|
|
51
|
+
return [f"{flag}={tokens[0]}"]
|
|
52
|
+
return [flag, *tokens]
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
def _flag_name(key: str) -> str:
|
|
56
|
+
"""Return the flag a key sets, whose value may be attached as `--flag=value`."""
|
|
57
|
+
return key.partition("=")[0]
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def get_server_arg(server_args: Mapping[str, str], flag: str) -> tuple[str, str] | None:
|
|
61
|
+
"""Return the key and value setting a flag under either spelling."""
|
|
62
|
+
for key, value in server_args.items():
|
|
63
|
+
if _flag_name(key) == flag:
|
|
64
|
+
return key, value
|
|
65
|
+
return None
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
def has_server_arg(server_args: Mapping[str, str], flag: str) -> bool:
|
|
69
|
+
"""Report whether a flag is set under either spelling, plain or attached."""
|
|
70
|
+
return get_server_arg(server_args, flag) is not None
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
def _merge_server_args(
|
|
74
|
+
defaults: Mapping[str, str], overrides: Mapping[str, str]
|
|
75
|
+
) -> dict[str, str]:
|
|
76
|
+
"""Overlay overrides on defaults by flag name, so an attached spelling wins."""
|
|
77
|
+
merged = {_flag_name(key): (key, value) for key, value in defaults.items()}
|
|
78
|
+
merged.update((_flag_name(key), (key, value)) for key, value in overrides.items())
|
|
79
|
+
return dict(merged.values())
|
|
80
|
+
|
|
81
|
+
|
|
39
82
|
class Endpoint(ABC):
|
|
40
83
|
"""A thing with a URL that you can start and stop."""
|
|
41
84
|
|
|
@@ -110,11 +153,15 @@ class SGLangEndpoint(Endpoint):
|
|
|
110
153
|
super().__init__(base_url=f"http://localhost:{sglang_port}")
|
|
111
154
|
|
|
112
155
|
self.worker_port = sglang_port
|
|
113
|
-
self.model_path = model_path
|
|
156
|
+
self.model_path = materialize_model_path(model_path)
|
|
114
157
|
self.tp = tp
|
|
115
158
|
self.ep = ep
|
|
116
159
|
self.dp = dp
|
|
117
|
-
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
|
+
)
|
|
118
165
|
self.load_format = load_format
|
|
119
166
|
self.nnodes = nnodes
|
|
120
167
|
self.node_rank = node_rank
|
|
@@ -123,9 +170,7 @@ class SGLangEndpoint(Endpoint):
|
|
|
123
170
|
self.disaggregation_mode = disaggregation_mode
|
|
124
171
|
self.prefill_bootstrap_port = prefill_bootstrap_port
|
|
125
172
|
self.launcher_module = launcher_module
|
|
126
|
-
self.extra_server_args = (
|
|
127
|
-
dict(extra_server_args) if extra_server_args else {}
|
|
128
|
-
)
|
|
173
|
+
self.extra_server_args = dict(extra_server_args) if extra_server_args else {}
|
|
129
174
|
self.health_timeout = health_timeout
|
|
130
175
|
self.health_poll_interval = health_poll_interval
|
|
131
176
|
self.health_request_timeout = health_request_timeout
|
|
@@ -140,28 +185,35 @@ class SGLangEndpoint(Endpoint):
|
|
|
140
185
|
self._bench_server: Optional[http.server.ThreadingHTTPServer] = None
|
|
141
186
|
|
|
142
187
|
if self.disaggregation_mode not in (None, "prefill", "decode"):
|
|
143
|
-
raise ValueError(
|
|
144
|
-
"disaggregation_mode must be None, 'prefill', or 'decode'"
|
|
145
|
-
)
|
|
188
|
+
raise ValueError("disaggregation_mode must be None, 'prefill', or 'decode'")
|
|
146
189
|
|
|
147
190
|
if os.environ.get(DUMMY_WEIGHTS_ENV) == "1":
|
|
148
|
-
if (
|
|
149
|
-
self.
|
|
150
|
-
and "--load-format" not in self.extra_server_args
|
|
191
|
+
if self.load_format is None and not has_server_arg(
|
|
192
|
+
self.extra_server_args, "--load-format"
|
|
151
193
|
):
|
|
152
194
|
self.load_format = "dummy"
|
|
153
|
-
|
|
154
|
-
|
|
195
|
+
if not has_server_arg(
|
|
196
|
+
self.extra_server_args, "--model-loader-extra-config"
|
|
197
|
+
):
|
|
198
|
+
self.extra_server_args["--model-loader-extra-config"] = "{}"
|
|
199
|
+
|
|
155
200
|
# Request logging enabled, log only metadata by default
|
|
156
201
|
if self.log_requests_level >= 0:
|
|
157
|
-
# Enable request logging
|
|
158
|
-
if "--log-requests"
|
|
159
|
-
self.extra_server_args
|
|
202
|
+
# Enable request logging
|
|
203
|
+
if not has_server_arg(self.extra_server_args, "--log-requests"):
|
|
204
|
+
self.extra_server_args["--log-requests"] = ""
|
|
160
205
|
# Set request logging level
|
|
161
|
-
|
|
162
|
-
|
|
206
|
+
level = get_server_arg(self.extra_server_args, "--log-requests-level")
|
|
207
|
+
if level is None:
|
|
208
|
+
self.extra_server_args["--log-requests-level"] = str(
|
|
209
|
+
self.log_requests_level
|
|
210
|
+
)
|
|
163
211
|
else:
|
|
164
|
-
|
|
212
|
+
key, value = level
|
|
213
|
+
print(
|
|
214
|
+
"[endpoint] --log-requests-level already set in server args: "
|
|
215
|
+
f"{value or key}"
|
|
216
|
+
)
|
|
165
217
|
|
|
166
218
|
# Add Modal-forwarded request headers to SGLang logged headers
|
|
167
219
|
sglang_log_request_headers = [
|
|
@@ -177,19 +229,21 @@ class SGLangEndpoint(Endpoint):
|
|
|
177
229
|
self.log_request_headers = ",".join(sglang_log_request_headers)
|
|
178
230
|
os.environ["SGLANG_LOG_REQUEST_HEADERS"] = self.log_request_headers
|
|
179
231
|
|
|
180
|
-
|
|
181
232
|
def _build_cmd(self) -> list[str]:
|
|
182
233
|
cmd = [
|
|
183
|
-
"python",
|
|
184
|
-
"
|
|
185
|
-
|
|
186
|
-
"--
|
|
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,
|
|
187
243
|
]
|
|
188
244
|
|
|
189
245
|
if self.speculative_model_path is not None:
|
|
190
|
-
cmd.extend(
|
|
191
|
-
["--speculative-draft-model-path", self.speculative_model_path]
|
|
192
|
-
)
|
|
246
|
+
cmd.extend(["--speculative-draft-model-path", self.speculative_model_path])
|
|
193
247
|
if self.load_format is not None:
|
|
194
248
|
cmd.extend(["--load-format", self.load_format])
|
|
195
249
|
if self.tp is not None:
|
|
@@ -214,19 +268,20 @@ class SGLangEndpoint(Endpoint):
|
|
|
214
268
|
raise ValueError("dist_init_host is required when nnodes > 1")
|
|
215
269
|
cmd.extend(
|
|
216
270
|
[
|
|
217
|
-
"--nnodes",
|
|
218
|
-
|
|
271
|
+
"--nnodes",
|
|
272
|
+
str(self.nnodes),
|
|
273
|
+
"--node-rank",
|
|
274
|
+
str(self.node_rank),
|
|
219
275
|
"--dist-init-addr",
|
|
220
276
|
f"{self.dist_init_host}:{self.dist_init_port}",
|
|
221
277
|
]
|
|
222
278
|
)
|
|
223
279
|
|
|
224
|
-
merged =
|
|
280
|
+
merged = _merge_server_args(
|
|
281
|
+
self.DEFAULT_OPERATIONAL_ARGS, self.extra_server_args
|
|
282
|
+
)
|
|
225
283
|
for key, value in merged.items():
|
|
226
|
-
|
|
227
|
-
cmd.append(key)
|
|
228
|
-
else:
|
|
229
|
-
cmd.extend([key, *value.split()])
|
|
284
|
+
cmd.extend(server_arg_tokens(key, value))
|
|
230
285
|
|
|
231
286
|
return cmd
|
|
232
287
|
|
|
@@ -287,9 +342,7 @@ class VLLMEndpoint(Endpoint):
|
|
|
287
342
|
super().__init__(base_url=f"http://localhost:{vllm_port}")
|
|
288
343
|
self.model = model
|
|
289
344
|
self.worker_port = vllm_port
|
|
290
|
-
self.extra_server_args = (
|
|
291
|
-
dict(extra_server_args) if extra_server_args else {}
|
|
292
|
-
)
|
|
345
|
+
self.extra_server_args = dict(extra_server_args) if extra_server_args else {}
|
|
293
346
|
self.health_timeout = health_timeout
|
|
294
347
|
self.health_poll_interval = health_poll_interval
|
|
295
348
|
self.health_request_timeout = health_request_timeout
|
|
@@ -297,20 +350,23 @@ class VLLMEndpoint(Endpoint):
|
|
|
297
350
|
self._bench_server: Optional[http.server.ThreadingHTTPServer] = None
|
|
298
351
|
|
|
299
352
|
if os.environ.get(DUMMY_WEIGHTS_ENV) == "1":
|
|
300
|
-
self.extra_server_args
|
|
353
|
+
if not has_server_arg(self.extra_server_args, "--load-format"):
|
|
354
|
+
self.extra_server_args["--load-format"] = "dummy"
|
|
301
355
|
|
|
302
356
|
def _build_cmd(self) -> list[str]:
|
|
303
357
|
cmd = [
|
|
304
|
-
"python",
|
|
305
|
-
"
|
|
306
|
-
"
|
|
307
|
-
"--
|
|
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,
|
|
308
367
|
]
|
|
309
368
|
for key, value in self.extra_server_args.items():
|
|
310
|
-
|
|
311
|
-
cmd.append(key)
|
|
312
|
-
else:
|
|
313
|
-
cmd.extend([key, *value.split()])
|
|
369
|
+
cmd.extend(server_arg_tokens(key, value))
|
|
314
370
|
return cmd
|
|
315
371
|
|
|
316
372
|
def start(self):
|
|
@@ -361,9 +417,7 @@ class TRTLLMEndpoint(Endpoint):
|
|
|
361
417
|
super().__init__(base_url=f"http://localhost:{worker_port}")
|
|
362
418
|
self.model = model
|
|
363
419
|
self.worker_port = worker_port
|
|
364
|
-
self.extra_server_args = (
|
|
365
|
-
dict(extra_server_args) if extra_server_args else {}
|
|
366
|
-
)
|
|
420
|
+
self.extra_server_args = dict(extra_server_args) if extra_server_args else {}
|
|
367
421
|
self.health_timeout = health_timeout
|
|
368
422
|
self.health_poll_interval = health_poll_interval
|
|
369
423
|
self.health_request_timeout = health_request_timeout
|
|
@@ -373,14 +427,13 @@ class TRTLLMEndpoint(Endpoint):
|
|
|
373
427
|
cmd = [
|
|
374
428
|
"trtllm-serve",
|
|
375
429
|
self.model,
|
|
376
|
-
"--host",
|
|
377
|
-
"
|
|
430
|
+
"--host",
|
|
431
|
+
"0.0.0.0",
|
|
432
|
+
"--port",
|
|
433
|
+
str(self.worker_port),
|
|
378
434
|
]
|
|
379
435
|
for key, value in self.extra_server_args.items():
|
|
380
|
-
|
|
381
|
-
cmd.append(key)
|
|
382
|
-
else:
|
|
383
|
-
cmd.extend([key, *value.split()])
|
|
436
|
+
cmd.extend(server_arg_tokens(key, value))
|
|
384
437
|
return cmd
|
|
385
438
|
|
|
386
439
|
def health_check(self) -> str | None:
|
|
@@ -442,18 +495,30 @@ class RouterEndpoint(Endpoint):
|
|
|
442
495
|
|
|
443
496
|
def _build_cmd(self) -> list[str]:
|
|
444
497
|
cmd = [
|
|
445
|
-
"python",
|
|
446
|
-
"
|
|
447
|
-
"
|
|
448
|
-
"--
|
|
449
|
-
"
|
|
450
|
-
"--
|
|
451
|
-
|
|
452
|
-
"--
|
|
453
|
-
"
|
|
454
|
-
"--
|
|
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",
|
|
455
519
|
"--disable-circuit-breaker",
|
|
456
|
-
"--request-timeout-secs",
|
|
520
|
+
"--request-timeout-secs",
|
|
521
|
+
"3600",
|
|
457
522
|
]
|
|
458
523
|
|
|
459
524
|
if self.api_key is not None:
|
|
@@ -465,9 +530,7 @@ class RouterEndpoint(Endpoint):
|
|
|
465
530
|
for role, node_ip in self.pd_config:
|
|
466
531
|
node_url = f"http://{node_ip}:{self.worker_port}"
|
|
467
532
|
if role == "prefill":
|
|
468
|
-
cmd.extend(
|
|
469
|
-
["--prefill", node_url, str(self.prefill_bootstrap_port)]
|
|
470
|
-
)
|
|
533
|
+
cmd.extend(["--prefill", node_url, str(self.prefill_bootstrap_port)])
|
|
471
534
|
elif role == "decode":
|
|
472
535
|
cmd.extend(["--decode", node_url])
|
|
473
536
|
elif role == "worker":
|
|
@@ -511,6 +574,7 @@ class RouterEndpoint(Endpoint):
|
|
|
511
574
|
# and this proxy listens on worker_port. POST /bench shells out to a benchmark
|
|
512
575
|
# task via run_bench(); every other path is forwarded to the upstream server.
|
|
513
576
|
|
|
577
|
+
|
|
514
578
|
def start_bench_proxy(
|
|
515
579
|
*,
|
|
516
580
|
listen_port: int,
|
|
@@ -521,17 +585,12 @@ def start_bench_proxy(
|
|
|
521
585
|
(_BenchProxyHandler,),
|
|
522
586
|
{"upstream_port": upstream_port},
|
|
523
587
|
)
|
|
524
|
-
server = http.server.ThreadingHTTPServer(
|
|
525
|
-
("0.0.0.0", listen_port), handler_cls
|
|
526
|
-
)
|
|
588
|
+
server = http.server.ThreadingHTTPServer(("0.0.0.0", listen_port), handler_cls)
|
|
527
589
|
thread = threading.Thread(
|
|
528
590
|
target=server.serve_forever, daemon=True, name="bench-proxy"
|
|
529
591
|
)
|
|
530
592
|
thread.start()
|
|
531
|
-
print(
|
|
532
|
-
f"[bench-proxy] listening on :{listen_port}, "
|
|
533
|
-
f"forwarding to :{upstream_port}"
|
|
534
|
-
)
|
|
593
|
+
print(f"[bench-proxy] listening on :{listen_port}, forwarding to :{upstream_port}")
|
|
535
594
|
return server
|
|
536
595
|
|
|
537
596
|
|
|
@@ -560,26 +619,17 @@ class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
|
|
|
560
619
|
try:
|
|
561
620
|
payload = self._read_json_body()
|
|
562
621
|
except (ValueError, json.JSONDecodeError) as exc:
|
|
563
|
-
return self._send_json(
|
|
564
|
-
400, {"ok": False, "error": f"bad body: {exc}"}
|
|
565
|
-
)
|
|
622
|
+
return self._send_json(400, {"ok": False, "error": f"bad body: {exc}"})
|
|
566
623
|
|
|
567
624
|
benchmark = payload.get("benchmark") or ""
|
|
568
625
|
args = payload.get("args") or []
|
|
569
|
-
target = (
|
|
570
|
-
payload.get("target")
|
|
571
|
-
or f"http://localhost:{self.upstream_port}"
|
|
572
|
-
)
|
|
626
|
+
target = payload.get("target") or f"http://localhost:{self.upstream_port}"
|
|
573
627
|
output_dir = payload.get("output_dir") or "/tmp/bench-output"
|
|
574
628
|
|
|
575
629
|
if not benchmark:
|
|
576
|
-
return self._send_json(
|
|
577
|
-
400, {"ok": False, "error": "benchmark is required"}
|
|
578
|
-
)
|
|
630
|
+
return self._send_json(400, {"ok": False, "error": "benchmark is required"})
|
|
579
631
|
if not isinstance(args, list):
|
|
580
|
-
return self._send_json(
|
|
581
|
-
400, {"ok": False, "error": "args must be a list"}
|
|
582
|
-
)
|
|
632
|
+
return self._send_json(400, {"ok": False, "error": "args must be a list"})
|
|
583
633
|
|
|
584
634
|
rc, body = run_bench(benchmark, list(args), target, output_dir)
|
|
585
635
|
status = 200 if rc == 0 else 500
|
|
@@ -630,9 +680,7 @@ class _BenchProxyHandler(http.server.BaseHTTPRequestHandler):
|
|
|
630
680
|
return json.loads(raw)
|
|
631
681
|
|
|
632
682
|
def _send_json(self, status: int, obj: Mapping[str, Any]):
|
|
633
|
-
self._send_raw(
|
|
634
|
-
status, "application/json", json.dumps(obj).encode()
|
|
635
|
-
)
|
|
683
|
+
self._send_raw(status, "application/json", json.dumps(obj).encode())
|
|
636
684
|
|
|
637
685
|
def _send_raw(self, status: int, content_type: str, body: bytes):
|
|
638
686
|
self.send_response(status)
|
|
@@ -657,18 +705,21 @@ def run_bench(
|
|
|
657
705
|
return 500, json.dumps(
|
|
658
706
|
{
|
|
659
707
|
"ok": False,
|
|
660
|
-
"error": (
|
|
661
|
-
"autoinference.benchmarks not importable in container"
|
|
662
|
-
),
|
|
708
|
+
"error": ("autoinference.benchmarks not importable in container"),
|
|
663
709
|
"volume_path": output_dir,
|
|
664
710
|
}
|
|
665
711
|
).encode()
|
|
666
712
|
|
|
667
713
|
cmd = [
|
|
668
|
-
sys.executable,
|
|
669
|
-
"
|
|
670
|
-
"
|
|
671
|
-
|
|
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],
|
|
672
723
|
]
|
|
673
724
|
env = {
|
|
674
725
|
**os.environ,
|
|
@@ -718,9 +769,7 @@ def wait_ready(
|
|
|
718
769
|
|
|
719
770
|
while time.time() < deadline:
|
|
720
771
|
try:
|
|
721
|
-
error = _health_check(
|
|
722
|
-
url, request_timeout=request_timeout, process=process
|
|
723
|
-
)
|
|
772
|
+
error = _health_check(url, request_timeout=request_timeout, process=process)
|
|
724
773
|
except subprocess.CalledProcessError as exc:
|
|
725
774
|
print(
|
|
726
775
|
f"[endpoint] !!! server process exited with code "
|
|
@@ -734,13 +783,11 @@ def wait_ready(
|
|
|
734
783
|
time.sleep(poll_interval)
|
|
735
784
|
|
|
736
785
|
print(
|
|
737
|
-
f"[endpoint] !!! health check timed out after {timeout}s; "
|
|
738
|
-
f"last={last_error}",
|
|
786
|
+
f"[endpoint] !!! health check timed out after {timeout}s; last={last_error}",
|
|
739
787
|
flush=True,
|
|
740
788
|
)
|
|
741
789
|
raise TimeoutError(
|
|
742
|
-
f"Health check timed out after {timeout}s for {url}. "
|
|
743
|
-
f"Last error: {last_error}"
|
|
790
|
+
f"Health check timed out after {timeout}s for {url}. Last error: {last_error}"
|
|
744
791
|
)
|
|
745
792
|
|
|
746
793
|
|
|
@@ -768,9 +815,7 @@ def warmup_chat_completions(
|
|
|
768
815
|
request_timeout=request_timeout,
|
|
769
816
|
max_attempts=max_attempts_per_request,
|
|
770
817
|
retry_delay=retry_delay,
|
|
771
|
-
description=(
|
|
772
|
-
f"warmup request {request_idx + 1}/{successful_requests}"
|
|
773
|
-
),
|
|
818
|
+
description=(f"warmup request {request_idx + 1}/{successful_requests}"),
|
|
774
819
|
)
|
|
775
820
|
|
|
776
821
|
|
|
@@ -792,9 +837,7 @@ def validate_embeddings_endpoint(
|
|
|
792
837
|
1 if all(isinstance(item, int) for item in inputs) else len(inputs)
|
|
793
838
|
)
|
|
794
839
|
else:
|
|
795
|
-
raise ValueError(
|
|
796
|
-
"embedding validation payload must contain non-empty input"
|
|
797
|
-
)
|
|
840
|
+
raise ValueError("embedding validation payload must contain non-empty input")
|
|
798
841
|
|
|
799
842
|
request_headers = {"Content-Type": "application/json"}
|
|
800
843
|
if headers:
|
|
@@ -857,10 +900,7 @@ def start_heartbeat_thread(
|
|
|
857
900
|
try:
|
|
858
901
|
error = health_check_fn()
|
|
859
902
|
except subprocess.CalledProcessError as exc:
|
|
860
|
-
print(
|
|
861
|
-
f"[heartbeat] server process exited with code "
|
|
862
|
-
f"{exc.returncode}"
|
|
863
|
-
)
|
|
903
|
+
print(f"[heartbeat] server process exited with code {exc.returncode}")
|
|
864
904
|
on_failure()
|
|
865
905
|
return
|
|
866
906
|
if error is None:
|
|
@@ -872,10 +912,7 @@ def start_heartbeat_thread(
|
|
|
872
912
|
f"({consecutive_failures}/{max_consecutive_failures})"
|
|
873
913
|
)
|
|
874
914
|
if consecutive_failures >= max_consecutive_failures:
|
|
875
|
-
print(
|
|
876
|
-
"[heartbeat] sustained health-check failure, "
|
|
877
|
-
"invoking on_failure"
|
|
878
|
-
)
|
|
915
|
+
print("[heartbeat] sustained health-check failure, invoking on_failure")
|
|
879
916
|
on_failure()
|
|
880
917
|
return
|
|
881
918
|
|
|
@@ -960,8 +997,7 @@ def _wait_ready_url(
|
|
|
960
997
|
last_error = error
|
|
961
998
|
time.sleep(poll_interval)
|
|
962
999
|
raise TimeoutError(
|
|
963
|
-
f"Timed out after {timeout}s waiting for {url}. "
|
|
964
|
-
f"Last error: {last_error}"
|
|
1000
|
+
f"Timed out after {timeout}s waiting for {url}. Last error: {last_error}"
|
|
965
1001
|
)
|
|
966
1002
|
|
|
967
1003
|
|
|
@@ -979,9 +1015,7 @@ def _post_json(
|
|
|
979
1015
|
timeout: float | None = None,
|
|
980
1016
|
) -> bytes:
|
|
981
1017
|
body = json.dumps(payload).encode("utf-8")
|
|
982
|
-
req = urllib.request.Request(
|
|
983
|
-
url, data=body, headers=dict(headers), method="POST"
|
|
984
|
-
)
|
|
1018
|
+
req = urllib.request.Request(url, data=body, headers=dict(headers), method="POST")
|
|
985
1019
|
with urllib.request.urlopen(req, timeout=timeout) as resp:
|
|
986
1020
|
return resp.read()
|
|
987
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
|