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.
@@ -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
- docs/*
244
- !docs/data-model.md
245
- !docs/dogfood-feedback.md
246
- !docs/main-architecture-explainer.md
247
- !docs/recipe-seeding.md
248
- !docs/sweep-learnings.md
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.4
1
+ Metadata-Version: 2.5
2
2
  Name: autoinference-utils
3
- Version: 0.2.5
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.5"
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 = ["autoinference_utils", "README.md", "pyproject.toml"]
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 = 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.load_format is None
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
- self.extra_server_args.setdefault("--model-loader-extra-config", "{}")
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" not in self.extra_server_args:
159
- self.extra_server_args.setdefault("--log-requests", "")
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
- if "--log-requests-level" not in self.extra_server_args:
162
- self.extra_server_args.setdefault("--log-requests-level", str(self.log_requests_level))
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
- print(f"[endpoint] --log-requests-level already set to {self.extra_server_args['--log-requests-level']} in server args")
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", "-m", self.launcher_module,
184
- "--host", "0.0.0.0",
185
- "--port", str(self.worker_port),
186
- "--model-path", self.model_path,
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", str(self.nnodes),
218
- "--node-rank", str(self.node_rank),
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 = {**self.DEFAULT_OPERATIONAL_ARGS, **self.extra_server_args}
280
+ merged = _merge_server_args(
281
+ self.DEFAULT_OPERATIONAL_ARGS, self.extra_server_args
282
+ )
225
283
  for key, value in merged.items():
226
- if value == "":
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.setdefault("--load-format", "dummy")
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", "-m", "vllm.entrypoints.openai.api_server",
305
- "--host", "0.0.0.0",
306
- "--port", str(self.worker_port),
307
- "--model", self.model,
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
- if value == "":
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", "0.0.0.0",
377
- "--port", str(self.worker_port),
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
- if value == "":
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", "-m", "sglang_router.launch_router",
446
- "--host", "0.0.0.0",
447
- "--port", str(self.router_port),
448
- "--prefill-policy", "cache_aware",
449
- "--decode-policy", "round_robin",
450
- "--max-concurrent-requests", "128",
451
- "--rate-limit-tokens-per-second", "0",
452
- "--queue-size", "0",
453
- "--health-check-timeout-secs", "600",
454
- "--log-level", "info",
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", "3600",
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, "-m", "invoke",
669
- "--search-root", search_root,
670
- "-c", "tasks",
671
- benchmark, *[str(a) for a in args],
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`