downshift-server 0.2.0__py3-none-any.whl
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.
- downshift/__init__.py +60 -0
- downshift/adapters/__init__.py +0 -0
- downshift/adapters/_flatten.py +40 -0
- downshift/adapters/base.py +47 -0
- downshift/adapters/generic.py +99 -0
- downshift/adapters/hf.py +95 -0
- downshift/adapters/pyg.py +120 -0
- downshift/adapters/registry.py +116 -0
- downshift/cli/__init__.py +0 -0
- downshift/cli/main.py +428 -0
- downshift/cli/render.py +174 -0
- downshift/export/__init__.py +0 -0
- downshift/export/capture.py +93 -0
- downshift/export/inputs.py +25 -0
- downshift/export/manifest.py +89 -0
- downshift/export/prevalidated.py +70 -0
- downshift/export/shapes.py +61 -0
- downshift/export/verdict.py +211 -0
- downshift/export/verify.py +162 -0
- downshift/loading.py +166 -0
- downshift/serve/__init__.py +0 -0
- downshift/serve/app.py +93 -0
- downshift/serve/backends.py +181 -0
- downshift/serve/engine.py +154 -0
- downshift/serve/middleware.py +24 -0
- downshift/serve/schemas.py +124 -0
- downshift/settings.py +69 -0
- downshift_server-0.2.0.dist-info/METADATA +265 -0
- downshift_server-0.2.0.dist-info/RECORD +32 -0
- downshift_server-0.2.0.dist-info/WHEEL +4 -0
- downshift_server-0.2.0.dist-info/entry_points.txt +7 -0
- downshift_server-0.2.0.dist-info/licenses/LICENSE +21 -0
downshift/cli/main.py
ADDED
|
@@ -0,0 +1,428 @@
|
|
|
1
|
+
"""downshift CLI. Each command loads, calls the library, and hands the result to render."""
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
import logging
|
|
5
|
+
import os
|
|
6
|
+
import re
|
|
7
|
+
import sys
|
|
8
|
+
from collections.abc import Iterator
|
|
9
|
+
from contextlib import contextmanager
|
|
10
|
+
from dataclasses import asdict, dataclass, replace
|
|
11
|
+
from enum import Enum
|
|
12
|
+
from pathlib import Path
|
|
13
|
+
from typing import Annotated
|
|
14
|
+
|
|
15
|
+
import typer
|
|
16
|
+
import uvicorn
|
|
17
|
+
from fastapi import FastAPI
|
|
18
|
+
|
|
19
|
+
import downshift
|
|
20
|
+
from downshift import __version__, settings
|
|
21
|
+
from downshift.cli import render
|
|
22
|
+
from downshift.export.manifest import manifest_path_for
|
|
23
|
+
from downshift.export.shapes import parse_dynamic_spec
|
|
24
|
+
from downshift.export.verdict import ExportVerdict
|
|
25
|
+
from downshift.loading import LoadedModel, LoadError, load_model
|
|
26
|
+
from downshift.serve.app import build_app
|
|
27
|
+
from downshift.serve.engine import BackendChoice, ServeOptions, ServingState, prepare_serving
|
|
28
|
+
|
|
29
|
+
EXIT_USAGE = 4 # bad model spec, bad option, unloadable file
|
|
30
|
+
EXIT_CRASH = 5
|
|
31
|
+
|
|
32
|
+
app = typer.Typer(
|
|
33
|
+
help="Check whether a PyTorch model survives ONNX export, then serve it.",
|
|
34
|
+
no_args_is_help=True,
|
|
35
|
+
add_completion=False,
|
|
36
|
+
)
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
class LogLevel(str, Enum):
|
|
40
|
+
debug = "debug"
|
|
41
|
+
info = "info"
|
|
42
|
+
warning = "warning"
|
|
43
|
+
error = "error"
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
class LogFormat(str, Enum):
|
|
47
|
+
text = "text"
|
|
48
|
+
json = "json"
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
ModelArg = Annotated[
|
|
52
|
+
str,
|
|
53
|
+
typer.Argument(metavar="MODEL", help="model.onnx | pkg.module:attr | weights.pt | org/repo | hf-repo-dir/"),
|
|
54
|
+
]
|
|
55
|
+
InputsOpt = Annotated[
|
|
56
|
+
str | None,
|
|
57
|
+
typer.Option(
|
|
58
|
+
"--inputs", metavar="pkg.module:fn", help="Example inputs: a tuple, or a factory for one"
|
|
59
|
+
),
|
|
60
|
+
]
|
|
61
|
+
ModelClassOpt = Annotated[
|
|
62
|
+
str | None,
|
|
63
|
+
typer.Option(
|
|
64
|
+
"--model-class", metavar="pkg.module:Class", help="Class to load a state dict into"
|
|
65
|
+
),
|
|
66
|
+
]
|
|
67
|
+
UnsafeLoadOpt = Annotated[
|
|
68
|
+
bool,
|
|
69
|
+
typer.Option(
|
|
70
|
+
"--unsafe-load", help="Allow torch.load(weights_only=False); runs code from the file"
|
|
71
|
+
),
|
|
72
|
+
]
|
|
73
|
+
AdapterOpt = Annotated[
|
|
74
|
+
str | None,
|
|
75
|
+
typer.Option(
|
|
76
|
+
"--adapter",
|
|
77
|
+
metavar="NAME|path/to/adapter.py[:attr]",
|
|
78
|
+
help="Model-family adapter: generic, pyg, hf, or your own adapter.py; default: detect",
|
|
79
|
+
),
|
|
80
|
+
]
|
|
81
|
+
SamplesOpt = Annotated[
|
|
82
|
+
int, typer.Option("-k", "--samples", min=1, help="Number of verification samples")
|
|
83
|
+
]
|
|
84
|
+
DynamicOpt = Annotated[
|
|
85
|
+
str | None,
|
|
86
|
+
typer.Option(
|
|
87
|
+
"--dynamic",
|
|
88
|
+
metavar="NAME:AXIS[,...]",
|
|
89
|
+
help='Dynamic axes, e.g. "x:0,edge_index:1". Default: axis 0 of every input.',
|
|
90
|
+
),
|
|
91
|
+
]
|
|
92
|
+
ReferenceOpt = Annotated[
|
|
93
|
+
str | None,
|
|
94
|
+
typer.Option("--reference", metavar="MODEL", help="PyTorch model to verify a .onnx against"),
|
|
95
|
+
]
|
|
96
|
+
IntraOpThreadsOpt = Annotated[
|
|
97
|
+
int,
|
|
98
|
+
typer.Option(
|
|
99
|
+
"--intra-op-threads", min=0, help="ORT threads within one op; 0 = let ONNX Runtime choose"
|
|
100
|
+
),
|
|
101
|
+
]
|
|
102
|
+
InterOpThreadsOpt = Annotated[
|
|
103
|
+
int,
|
|
104
|
+
typer.Option(
|
|
105
|
+
"--inter-op-threads", min=0, help="ORT threads across ops; 0 = let ONNX Runtime choose"
|
|
106
|
+
),
|
|
107
|
+
]
|
|
108
|
+
WorkersOpt = Annotated[
|
|
109
|
+
int,
|
|
110
|
+
typer.Option(
|
|
111
|
+
"--workers",
|
|
112
|
+
min=1,
|
|
113
|
+
help="Uvicorn worker processes; each independently loads/exports/warms the model",
|
|
114
|
+
),
|
|
115
|
+
]
|
|
116
|
+
JsonOpt = Annotated[
|
|
117
|
+
bool, typer.Option("--json", help="Print the verdict as JSON and nothing else")
|
|
118
|
+
]
|
|
119
|
+
LogLevelOpt = Annotated[LogLevel, typer.Option("--log-level")]
|
|
120
|
+
LogFormatOpt = Annotated[LogFormat, typer.Option("--log-format")]
|
|
121
|
+
|
|
122
|
+
|
|
123
|
+
class _JsonFormatter(logging.Formatter):
|
|
124
|
+
def format(self, record: logging.LogRecord) -> str:
|
|
125
|
+
payload = {
|
|
126
|
+
"time": self.formatTime(record, "%Y-%m-%dT%H:%M:%S"),
|
|
127
|
+
"level": record.levelname,
|
|
128
|
+
"logger": record.name,
|
|
129
|
+
"message": record.getMessage(),
|
|
130
|
+
}
|
|
131
|
+
if record.exc_info:
|
|
132
|
+
payload["exc_info"] = self.formatException(record.exc_info)
|
|
133
|
+
return json.dumps(payload)
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
def _setup_logging(level: LogLevel, fmt: LogFormat) -> None:
|
|
137
|
+
handler = logging.StreamHandler(sys.stderr) # keep stdout clean for --json
|
|
138
|
+
if fmt is LogFormat.json:
|
|
139
|
+
handler.setFormatter(_JsonFormatter())
|
|
140
|
+
else:
|
|
141
|
+
handler.setFormatter(logging.Formatter("%(levelname)s %(name)s: %(message)s"))
|
|
142
|
+
logging.basicConfig(level=level.value.upper(), handlers=[handler], force=True)
|
|
143
|
+
# The ONNX optimizer passes log every rewrite at INFO. Only show them when debugging.
|
|
144
|
+
if level is not LogLevel.debug:
|
|
145
|
+
for name in ("onnxscript", "onnx_ir"):
|
|
146
|
+
logging.getLogger(name).setLevel(logging.WARNING)
|
|
147
|
+
|
|
148
|
+
|
|
149
|
+
@contextmanager
|
|
150
|
+
def _exit_on_error(debug: bool) -> Iterator[None]:
|
|
151
|
+
"""User errors exit 4, anything else exits 5. typer.Exit passes through untouched."""
|
|
152
|
+
try:
|
|
153
|
+
yield
|
|
154
|
+
except typer.Exit:
|
|
155
|
+
raise
|
|
156
|
+
except ValueError as exc: # LoadError, bad --dynamic, bad backend combination
|
|
157
|
+
render.error(str(exc))
|
|
158
|
+
raise typer.Exit(EXIT_USAGE) from exc
|
|
159
|
+
except Exception as exc:
|
|
160
|
+
if debug:
|
|
161
|
+
render.print_traceback()
|
|
162
|
+
render.error(f"{type(exc).__name__}: {exc}")
|
|
163
|
+
raise typer.Exit(EXIT_CRASH) from exc
|
|
164
|
+
|
|
165
|
+
|
|
166
|
+
def _load(spec: str, inputs: str | None, model_class: str | None, unsafe_load: bool) -> LoadedModel:
|
|
167
|
+
if unsafe_load:
|
|
168
|
+
render.warn(
|
|
169
|
+
f"--unsafe-load: torch.load(weights_only=False) on {spec}; arbitrary code may run"
|
|
170
|
+
)
|
|
171
|
+
return load_model(spec, inputs=inputs, model_class=model_class, unsafe_load=unsafe_load)
|
|
172
|
+
|
|
173
|
+
|
|
174
|
+
_IMPORT_SPEC = re.compile(r"^[A-Za-z_][\w.]*:[A-Za-z_]\w*$")
|
|
175
|
+
|
|
176
|
+
|
|
177
|
+
def slug(spec: str) -> str:
|
|
178
|
+
"""tests.models.clean_mlp:make_model -> clean_mlp; ./gat_v3.pt -> gat_v3; org/repo -> repo."""
|
|
179
|
+
if _IMPORT_SPEC.match(spec):
|
|
180
|
+
return spec.partition(":")[0].rsplit(".", 1)[-1]
|
|
181
|
+
return Path(spec).stem
|
|
182
|
+
|
|
183
|
+
|
|
184
|
+
def _emit(
|
|
185
|
+
verdict: ExportVerdict, model_name: str, as_json: bool, extra: dict | None = None
|
|
186
|
+
) -> None:
|
|
187
|
+
if as_json:
|
|
188
|
+
typer.echo(json.dumps(verdict.to_dict() | (extra or {}), indent=2))
|
|
189
|
+
else:
|
|
190
|
+
render.print_verdict(verdict, model_name)
|
|
191
|
+
|
|
192
|
+
|
|
193
|
+
@app.command("check")
|
|
194
|
+
def check_cmd(
|
|
195
|
+
model: ModelArg,
|
|
196
|
+
json_out: JsonOpt = False,
|
|
197
|
+
reference: ReferenceOpt = None,
|
|
198
|
+
inputs: InputsOpt = None,
|
|
199
|
+
model_class: ModelClassOpt = None,
|
|
200
|
+
unsafe_load: UnsafeLoadOpt = False,
|
|
201
|
+
adapter: AdapterOpt = None,
|
|
202
|
+
k: SamplesOpt = settings.SAMPLES,
|
|
203
|
+
dynamic: DynamicOpt = None,
|
|
204
|
+
log_level: LogLevelOpt = LogLevel.warning,
|
|
205
|
+
log_format: LogFormatOpt = LogFormat.text,
|
|
206
|
+
) -> None:
|
|
207
|
+
"""Export in memory and verify numerics. Exit 0 CLEAN, 1 FAILED, 2 DEGRADED, 3 UNVERIFIED."""
|
|
208
|
+
_setup_logging(log_level, log_format)
|
|
209
|
+
with _exit_on_error(log_level is LogLevel.debug):
|
|
210
|
+
loaded = _load(model, inputs, model_class, unsafe_load)
|
|
211
|
+
dynamic_spec = parse_dynamic_spec(dynamic) if dynamic else None
|
|
212
|
+
if loaded.onnx_path is not None:
|
|
213
|
+
ref = _load(reference, inputs, model_class, unsafe_load) if reference else None
|
|
214
|
+
verdict = downshift.intake(
|
|
215
|
+
loaded.onnx_path,
|
|
216
|
+
ref.model if ref else None,
|
|
217
|
+
ref.example_inputs if ref else None,
|
|
218
|
+
adapter or (ref.adapter_hint if ref else None),
|
|
219
|
+
k=k,
|
|
220
|
+
dynamic=dynamic_spec,
|
|
221
|
+
)
|
|
222
|
+
else:
|
|
223
|
+
assert loaded.model is not None
|
|
224
|
+
verdict = downshift.check(
|
|
225
|
+
loaded.model,
|
|
226
|
+
loaded.example_inputs,
|
|
227
|
+
k=k,
|
|
228
|
+
adapter=adapter or loaded.adapter_hint,
|
|
229
|
+
dynamic=dynamic_spec,
|
|
230
|
+
)
|
|
231
|
+
_emit(verdict, model, json_out)
|
|
232
|
+
raise typer.Exit(verdict.exit_code)
|
|
233
|
+
|
|
234
|
+
|
|
235
|
+
@app.command("export")
|
|
236
|
+
def export_cmd(
|
|
237
|
+
model: ModelArg,
|
|
238
|
+
output: Annotated[Path, typer.Option("-o", "--output", help="Output directory")],
|
|
239
|
+
name: Annotated[
|
|
240
|
+
str | None, typer.Option("--name", help="Artifact stem; default: model slug")
|
|
241
|
+
] = None,
|
|
242
|
+
fp16: Annotated[
|
|
243
|
+
bool, typer.Option("--fp16", help="Cast the model to fp16 before export")
|
|
244
|
+
] = False,
|
|
245
|
+
no_verify: Annotated[
|
|
246
|
+
bool, typer.Option("--no-verify", help="Skip numerics; the verdict is UNVERIFIED")
|
|
247
|
+
] = False,
|
|
248
|
+
json_out: JsonOpt = False,
|
|
249
|
+
inputs: InputsOpt = None,
|
|
250
|
+
model_class: ModelClassOpt = None,
|
|
251
|
+
unsafe_load: UnsafeLoadOpt = False,
|
|
252
|
+
adapter: AdapterOpt = None,
|
|
253
|
+
k: SamplesOpt = settings.SAMPLES,
|
|
254
|
+
dynamic: DynamicOpt = None,
|
|
255
|
+
log_level: LogLevelOpt = LogLevel.warning,
|
|
256
|
+
log_format: LogFormatOpt = LogFormat.text,
|
|
257
|
+
) -> None:
|
|
258
|
+
"""Export to DIR/NAME.onnx with a NAME.manifest.json sidecar. Nothing is written if FAILED."""
|
|
259
|
+
_setup_logging(log_level, log_format)
|
|
260
|
+
with _exit_on_error(log_level is LogLevel.debug):
|
|
261
|
+
loaded = _load(model, inputs, model_class, unsafe_load)
|
|
262
|
+
if loaded.model is None:
|
|
263
|
+
raise LoadError(f"{model} is already ONNX; export needs a PyTorch model")
|
|
264
|
+
onnx_path = output / f"{name or slug(model)}.onnx"
|
|
265
|
+
if no_verify:
|
|
266
|
+
render.warn("--no-verify: the graph is saved without checking its numerics")
|
|
267
|
+
verdict = downshift.export(
|
|
268
|
+
loaded.model,
|
|
269
|
+
onnx_path,
|
|
270
|
+
loaded.example_inputs,
|
|
271
|
+
k=k,
|
|
272
|
+
adapter=adapter or loaded.adapter_hint,
|
|
273
|
+
dynamic=parse_dynamic_spec(dynamic) if dynamic else None,
|
|
274
|
+
fp16=fp16,
|
|
275
|
+
source_path=loaded.source_path,
|
|
276
|
+
verify_numerics=not no_verify,
|
|
277
|
+
)
|
|
278
|
+
manifest = manifest_path_for(onnx_path) if verdict.onnx_path else None
|
|
279
|
+
_emit(verdict, model, json_out, {"manifest_path": str(manifest) if manifest else None})
|
|
280
|
+
if not json_out:
|
|
281
|
+
render.print_artifacts(verdict.onnx_path, manifest)
|
|
282
|
+
raise typer.Exit(verdict.exit_code)
|
|
283
|
+
|
|
284
|
+
|
|
285
|
+
@dataclass
|
|
286
|
+
class ServeArgs:
|
|
287
|
+
"""Everything needed to rebuild a ServingState + FastAPI app from scratch. Plain JSON-able
|
|
288
|
+
types only: a multi-worker run ships this to each worker process via an env var."""
|
|
289
|
+
|
|
290
|
+
model: str
|
|
291
|
+
inputs: str | None
|
|
292
|
+
model_class: str | None
|
|
293
|
+
unsafe_load: bool
|
|
294
|
+
adapter: str | None
|
|
295
|
+
k: int
|
|
296
|
+
dynamic: str | None
|
|
297
|
+
reference: str | None
|
|
298
|
+
middleware: list[str] | None
|
|
299
|
+
backend: str
|
|
300
|
+
force_onnx: bool
|
|
301
|
+
device: str
|
|
302
|
+
warmup: int
|
|
303
|
+
intra_op_threads: int
|
|
304
|
+
inter_op_threads: int
|
|
305
|
+
log_level: str
|
|
306
|
+
log_format: str
|
|
307
|
+
|
|
308
|
+
|
|
309
|
+
_SERVE_ARGS_ENV = "_DOWNSHIFT_SERVE_ARGS"
|
|
310
|
+
|
|
311
|
+
|
|
312
|
+
def _build_serving_app(args: ServeArgs) -> tuple[ServingState, FastAPI]:
|
|
313
|
+
loaded = _load(args.model, args.inputs, args.model_class, args.unsafe_load)
|
|
314
|
+
ref = (
|
|
315
|
+
_load(args.reference, args.inputs, args.model_class, args.unsafe_load)
|
|
316
|
+
if args.reference
|
|
317
|
+
else None
|
|
318
|
+
)
|
|
319
|
+
opts = ServeOptions(
|
|
320
|
+
backend=BackendChoice(args.backend),
|
|
321
|
+
force_onnx=args.force_onnx,
|
|
322
|
+
device=args.device,
|
|
323
|
+
warmup=args.warmup,
|
|
324
|
+
k=args.k,
|
|
325
|
+
adapter=args.adapter,
|
|
326
|
+
dynamic=parse_dynamic_spec(args.dynamic) if args.dynamic else None,
|
|
327
|
+
intra_op_threads=args.intra_op_threads,
|
|
328
|
+
inter_op_threads=args.inter_op_threads,
|
|
329
|
+
)
|
|
330
|
+
state = prepare_serving(loaded, opts, ref)
|
|
331
|
+
api = build_app(state, tuple(args.middleware or ()))
|
|
332
|
+
return state, api
|
|
333
|
+
|
|
334
|
+
|
|
335
|
+
def _serve_app_factory() -> FastAPI:
|
|
336
|
+
"""Import-string target for uvicorn's multi-worker mode (`downshift.cli.main:_serve_app_factory`).
|
|
337
|
+
Each worker process calls this on its own, independently reloading/re-exporting/re-warming
|
|
338
|
+
the model from the args the parent process serialized into _SERVE_ARGS_ENV."""
|
|
339
|
+
args = ServeArgs(**json.loads(os.environ[_SERVE_ARGS_ENV]))
|
|
340
|
+
_setup_logging(LogLevel(args.log_level), LogFormat(args.log_format))
|
|
341
|
+
_, api = _build_serving_app(args)
|
|
342
|
+
return api
|
|
343
|
+
|
|
344
|
+
|
|
345
|
+
@app.command("serve")
|
|
346
|
+
def serve_cmd(
|
|
347
|
+
model: ModelArg,
|
|
348
|
+
host: Annotated[str, typer.Option("--host")] = settings.HOST,
|
|
349
|
+
port: Annotated[int, typer.Option("--port")] = settings.PORT,
|
|
350
|
+
backend: Annotated[BackendChoice, typer.Option("--backend")] = BackendChoice(settings.BACKEND),
|
|
351
|
+
force_onnx: Annotated[
|
|
352
|
+
bool, typer.Option("--force-onnx", help="Serve a DEGRADED graph via ONNX Runtime anyway")
|
|
353
|
+
] = False,
|
|
354
|
+
device: Annotated[str, typer.Option("--device", help="auto | cpu | cuda")] = settings.DEVICE,
|
|
355
|
+
warmup: Annotated[
|
|
356
|
+
int, typer.Option("--warmup", min=0, help="Warm-up inferences before /ready flips")
|
|
357
|
+
] = settings.WARMUP,
|
|
358
|
+
reference: ReferenceOpt = None,
|
|
359
|
+
middleware: Annotated[
|
|
360
|
+
list[str] | None,
|
|
361
|
+
typer.Option("--middleware", metavar="pkg.module:Attr", help="Middleware to attach; repeatable"),
|
|
362
|
+
] = None,
|
|
363
|
+
inputs: InputsOpt = None,
|
|
364
|
+
model_class: ModelClassOpt = None,
|
|
365
|
+
unsafe_load: UnsafeLoadOpt = False,
|
|
366
|
+
adapter: AdapterOpt = None,
|
|
367
|
+
k: SamplesOpt = settings.SAMPLES,
|
|
368
|
+
dynamic: DynamicOpt = None,
|
|
369
|
+
intra_op_threads: IntraOpThreadsOpt = settings.INTRA_OP_THREADS,
|
|
370
|
+
inter_op_threads: InterOpThreadsOpt = settings.INTER_OP_THREADS,
|
|
371
|
+
workers: WorkersOpt = settings.WORKERS,
|
|
372
|
+
log_level: LogLevelOpt = LogLevel.info,
|
|
373
|
+
log_format: LogFormatOpt = LogFormat.text,
|
|
374
|
+
) -> None:
|
|
375
|
+
"""Check the model, pick a backend from the verdict, and serve it over HTTP."""
|
|
376
|
+
_setup_logging(log_level, log_format)
|
|
377
|
+
with _exit_on_error(log_level is LogLevel.debug):
|
|
378
|
+
args = ServeArgs(
|
|
379
|
+
model=model,
|
|
380
|
+
inputs=inputs,
|
|
381
|
+
model_class=model_class,
|
|
382
|
+
unsafe_load=unsafe_load,
|
|
383
|
+
adapter=adapter,
|
|
384
|
+
k=k,
|
|
385
|
+
dynamic=dynamic,
|
|
386
|
+
reference=reference,
|
|
387
|
+
middleware=list(middleware) if middleware else None,
|
|
388
|
+
backend=backend.value,
|
|
389
|
+
force_onnx=force_onnx,
|
|
390
|
+
device=device,
|
|
391
|
+
warmup=warmup,
|
|
392
|
+
intra_op_threads=intra_op_threads,
|
|
393
|
+
inter_op_threads=inter_op_threads,
|
|
394
|
+
log_level=log_level.value,
|
|
395
|
+
log_format=log_format.value,
|
|
396
|
+
)
|
|
397
|
+
if workers <= 1:
|
|
398
|
+
state, api = _build_serving_app(args)
|
|
399
|
+
render.print_banner(state, host, port)
|
|
400
|
+
uvicorn.run(api, host=host, port=port, log_level=log_level.value)
|
|
401
|
+
else:
|
|
402
|
+
# Only for the banner/fail-fast check: each of the N workers rebuilds its own
|
|
403
|
+
# backend anyway, so this throwaway copy skips warmup, it'll never serve traffic.
|
|
404
|
+
state, _ = _build_serving_app(replace(args, warmup=0))
|
|
405
|
+
render.print_banner(state, host, port)
|
|
406
|
+
render.warn(
|
|
407
|
+
f"--workers {workers}: each worker independently reloads, re-exports, and "
|
|
408
|
+
"re-warms the model (memory and startup time scale with this number)"
|
|
409
|
+
)
|
|
410
|
+
os.environ[_SERVE_ARGS_ENV] = json.dumps(asdict(args))
|
|
411
|
+
uvicorn.run(
|
|
412
|
+
"downshift.cli.main:_serve_app_factory",
|
|
413
|
+
host=host,
|
|
414
|
+
port=port,
|
|
415
|
+
workers=workers,
|
|
416
|
+
log_level=log_level.value,
|
|
417
|
+
factory=True,
|
|
418
|
+
)
|
|
419
|
+
|
|
420
|
+
|
|
421
|
+
@app.command()
|
|
422
|
+
def version() -> None:
|
|
423
|
+
"""Print the version."""
|
|
424
|
+
typer.echo(f"downshift v{__version__}")
|
|
425
|
+
|
|
426
|
+
|
|
427
|
+
if __name__ == "__main__": # pragma: no cover
|
|
428
|
+
app()
|
downshift/cli/render.py
ADDED
|
@@ -0,0 +1,174 @@
|
|
|
1
|
+
"""All rich output for the CLI lives here. Commands hand over objects; this module prints."""
|
|
2
|
+
|
|
3
|
+
from pathlib import Path
|
|
4
|
+
|
|
5
|
+
from rich import box
|
|
6
|
+
from rich.console import Console
|
|
7
|
+
from rich.markup import escape
|
|
8
|
+
from rich.panel import Panel
|
|
9
|
+
from rich.table import Table
|
|
10
|
+
from rich.text import Text
|
|
11
|
+
|
|
12
|
+
from downshift import __version__
|
|
13
|
+
from downshift.export.verdict import ExportVerdict
|
|
14
|
+
from downshift.serve.engine import ServingState
|
|
15
|
+
|
|
16
|
+
console = Console()
|
|
17
|
+
err_console = Console(stderr=True)
|
|
18
|
+
|
|
19
|
+
STATUS_STYLE = {
|
|
20
|
+
"CLEAN": "bold green",
|
|
21
|
+
"DEGRADED": "bold yellow",
|
|
22
|
+
"FAILED": "bold red",
|
|
23
|
+
"UNVERIFIED": "bold magenta",
|
|
24
|
+
}
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def _sym(utf: str, ascii_: str) -> str:
|
|
28
|
+
"""Unicode glyph unless the console can't encode it (cp1252 Windows pipes)."""
|
|
29
|
+
return ascii_ if console.options.ascii_only else utf
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def _status_text(verdict: ExportVerdict) -> Text:
|
|
33
|
+
text = Text(verdict.status, style=STATUS_STYLE[verdict.status])
|
|
34
|
+
details = [verdict.capture_strategy] if verdict.capture_strategy else []
|
|
35
|
+
if verdict.opset is not None:
|
|
36
|
+
details.append(f"opset {verdict.opset}")
|
|
37
|
+
if details:
|
|
38
|
+
text.append(f" ({', '.join(details)})", style="dim")
|
|
39
|
+
return text
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def _numerics_text(verdict: ExportVerdict) -> Text:
|
|
43
|
+
n = verdict.numerics
|
|
44
|
+
if n is None:
|
|
45
|
+
return Text("not checked", style="dim")
|
|
46
|
+
text = Text(f"max abs err {n.max_abs_err:.2e} over {n.samples_tested} samples ")
|
|
47
|
+
if n.passed:
|
|
48
|
+
text.append(_sym("✓", "OK"), style="bold green")
|
|
49
|
+
else:
|
|
50
|
+
text.append(f"{_sym('✗', 'X')} {n.failures}/{n.samples_tested} failed", style="bold red")
|
|
51
|
+
return text
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def _dynamic_text(verdict: ExportVerdict) -> str:
|
|
55
|
+
if not verdict.dynamic_dims:
|
|
56
|
+
return _sym("—", "-")
|
|
57
|
+
return ", ".join(
|
|
58
|
+
f"{name}[{axis}]" for name, axes in verdict.dynamic_dims.items() for axis in axes
|
|
59
|
+
)
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
def _shape_text(verdict: ExportVerdict) -> str:
|
|
63
|
+
if verdict.shape_generalization is None:
|
|
64
|
+
return _sym("—", "-")
|
|
65
|
+
return "yes" if verdict.shape_generalization else "no"
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
def print_verdict(verdict: ExportVerdict, model_name: str) -> None:
|
|
69
|
+
table = Table(show_header=False, box=box.ROUNDED, border_style=STATUS_STYLE[verdict.status])
|
|
70
|
+
table.add_column(style="bold", no_wrap=True)
|
|
71
|
+
table.add_column()
|
|
72
|
+
table.add_row("Model", escape(model_name))
|
|
73
|
+
table.add_row("Family", verdict.model_family)
|
|
74
|
+
table.add_row("Export", _status_text(verdict))
|
|
75
|
+
table.add_row("Numerics", _numerics_text(verdict))
|
|
76
|
+
table.add_row("Shape-general", _shape_text(verdict))
|
|
77
|
+
table.add_row("Dynamic dims", _dynamic_text(verdict))
|
|
78
|
+
if verdict.unsupported_ops:
|
|
79
|
+
table.add_row("Unsupported ops", Text(", ".join(verdict.unsupported_ops), style="red"))
|
|
80
|
+
if verdict.warnings:
|
|
81
|
+
table.add_row("Warnings", Text("\n".join(verdict.warnings), style="yellow"))
|
|
82
|
+
table.add_row("Backend", verdict.recommended_backend)
|
|
83
|
+
table.add_row("Reason", escape(verdict.reason))
|
|
84
|
+
console.print(table)
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
def print_artifacts(onnx_path: Path | None, manifest_path: Path | None) -> None:
|
|
88
|
+
if onnx_path is None:
|
|
89
|
+
console.print("[red]nothing written[/]: the export failed")
|
|
90
|
+
return
|
|
91
|
+
console.print(f"[bold]Wrote[/] {escape(str(onnx_path))}")
|
|
92
|
+
if manifest_path is not None:
|
|
93
|
+
console.print(f"[bold]Manifest[/] {escape(str(manifest_path))}")
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
def _backend_text(state: ServingState) -> Text:
|
|
97
|
+
meta = state.backend.metadata()
|
|
98
|
+
label = "torch (eager)" if meta.name == "torch" else meta.name
|
|
99
|
+
text = Text(f"{label} {_sym('·', '|')} {meta.device}")
|
|
100
|
+
arrow = _sym("←", "<-")
|
|
101
|
+
if state.forced_onnx:
|
|
102
|
+
text.append(f" {arrow} --force-onnx", style="yellow")
|
|
103
|
+
elif state.backend_auto_selected:
|
|
104
|
+
text.append(f" {arrow} auto-selected", style="dim")
|
|
105
|
+
return text
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
def print_banner(state: ServingState, host: str, port: int) -> None:
|
|
109
|
+
"""Boot banner for `serve`. Says what is served, how it was judged, and where it listens."""
|
|
110
|
+
verdict = state.verdict
|
|
111
|
+
sub = _sym("└", "\\")
|
|
112
|
+
grid = Table.grid(padding=(0, 3))
|
|
113
|
+
grid.add_column(style="bold cyan", no_wrap=True)
|
|
114
|
+
grid.add_column()
|
|
115
|
+
|
|
116
|
+
grid.add_row("Model", escape(state.source))
|
|
117
|
+
grid.add_row("Family", verdict.model_family)
|
|
118
|
+
|
|
119
|
+
unverified_onnx = verdict.status == "UNVERIFIED" and verdict.prepared is None
|
|
120
|
+
if unverified_onnx:
|
|
121
|
+
grid.add_row("Verdict", _status_text(verdict) + Text(" - no reference model supplied"))
|
|
122
|
+
grid.add_row("", Text(f"{sub} served as-is; numerics were never checked", style="dim"))
|
|
123
|
+
grid.add_row("Tip", "pass --reference <model> to verify")
|
|
124
|
+
else:
|
|
125
|
+
grid.add_row("Verdict", _status_text(verdict))
|
|
126
|
+
if verdict.status in ("FAILED", "UNVERIFIED"):
|
|
127
|
+
grid.add_row("", Text(f"{sub} {verdict.reason}", style="dim"))
|
|
128
|
+
|
|
129
|
+
grid.add_row("Numerics", _numerics_text(verdict))
|
|
130
|
+
|
|
131
|
+
if verdict.status == "DEGRADED":
|
|
132
|
+
n = verdict.numerics
|
|
133
|
+
if n is None:
|
|
134
|
+
detail = verdict.reason
|
|
135
|
+
else:
|
|
136
|
+
detail = (
|
|
137
|
+
f"numerics diverge on {n.failures}/{n.samples_tested} samples "
|
|
138
|
+
f"(max abs err {n.max_abs_err:.2e})"
|
|
139
|
+
)
|
|
140
|
+
grid.add_row("", Text(f"{_sym('⚠', '!')} {detail}", style="yellow"))
|
|
141
|
+
if state.backend.name != "onnxruntime":
|
|
142
|
+
grid.add_row("Override", "--force-onnx to serve the ONNX graph anyway")
|
|
143
|
+
|
|
144
|
+
grid.add_row("Backend", _backend_text(state))
|
|
145
|
+
for note in state.notes: # backend-selection notes from the engine
|
|
146
|
+
style = "bold red" if "outputs may be wrong" in note else "yellow"
|
|
147
|
+
grid.add_row("", Text(f"{_sym('⚠', '!')} {note}", style=style))
|
|
148
|
+
grid.add_row("Dynamic dims", _dynamic_text(verdict))
|
|
149
|
+
for warning in verdict.warnings:
|
|
150
|
+
grid.add_row("", Text(f"{_sym('⚠', '!')} {warning}", style="yellow"))
|
|
151
|
+
grid.add_row("Endpoint", f"http://{host}:{port}")
|
|
152
|
+
|
|
153
|
+
console.print(
|
|
154
|
+
Panel(
|
|
155
|
+
grid,
|
|
156
|
+
title=f"downshift v{__version__}",
|
|
157
|
+
title_align="left",
|
|
158
|
+
border_style=STATUS_STYLE[verdict.status],
|
|
159
|
+
expand=False,
|
|
160
|
+
padding=(1, 2),
|
|
161
|
+
)
|
|
162
|
+
)
|
|
163
|
+
|
|
164
|
+
|
|
165
|
+
def warn(msg: str) -> None:
|
|
166
|
+
err_console.print(f"[bold yellow]warning:[/] {escape(msg)}")
|
|
167
|
+
|
|
168
|
+
|
|
169
|
+
def error(msg: str) -> None:
|
|
170
|
+
err_console.print(f"[bold red]error:[/] {escape(msg)}")
|
|
171
|
+
|
|
172
|
+
|
|
173
|
+
def print_traceback() -> None:
|
|
174
|
+
err_console.print_exception()
|
|
File without changes
|
|
@@ -0,0 +1,93 @@
|
|
|
1
|
+
"""Drive torch.export + torch.onnx.export and say which strategy worked.
|
|
2
|
+
|
|
3
|
+
We run torch.export ourselves (strict=False, then strict=True) so the winning strategy is
|
|
4
|
+
a value we return, not something scraped from console output. The ONNX translation is
|
|
5
|
+
still entirely torch.onnx.export's; we only hand it the ExportedProgram.
|
|
6
|
+
|
|
7
|
+
torch 2.14 note: the dynamo exporter documents only the two strict modes. There's no
|
|
8
|
+
draft_export step or TorchScript fallback any more.
|
|
9
|
+
"""
|
|
10
|
+
|
|
11
|
+
import contextlib
|
|
12
|
+
import io
|
|
13
|
+
import logging
|
|
14
|
+
from dataclasses import dataclass, field
|
|
15
|
+
|
|
16
|
+
import torch
|
|
17
|
+
|
|
18
|
+
_STRATEGIES: tuple[tuple[str, bool], ...] = (("strict=False", False), ("strict=True", True))
|
|
19
|
+
|
|
20
|
+
# torch.onnx logs a warning per missing torchvision op on every export. Not actionable.
|
|
21
|
+
logging.getLogger("torch.onnx._internal.exporter._registration").setLevel(logging.ERROR)
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
@dataclass
|
|
25
|
+
class CaptureResult:
|
|
26
|
+
success: bool
|
|
27
|
+
capture_strategy: str | None # a _STRATEGIES name; None when nothing traced
|
|
28
|
+
onnx_program: "torch.onnx.ONNXProgram | None" = None
|
|
29
|
+
opset: int | None = None
|
|
30
|
+
op_types: list[str] = field(default_factory=list)
|
|
31
|
+
exception: Exception | None = None
|
|
32
|
+
stderr: str = "" # whatever torch printed while we tried; useful at debug level
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def capture(
|
|
36
|
+
model: torch.nn.Module,
|
|
37
|
+
example_inputs: tuple,
|
|
38
|
+
dynamic_shapes: tuple | None = None,
|
|
39
|
+
) -> CaptureResult:
|
|
40
|
+
exported_program = None
|
|
41
|
+
strategy_used: str | None = None
|
|
42
|
+
last_exception: Exception | None = None
|
|
43
|
+
# torch prints whole FX graphs straight to stderr when a data-dependent guard fails.
|
|
44
|
+
# Keep that out of the user's terminal; the exception message is what matters.
|
|
45
|
+
captured = io.StringIO()
|
|
46
|
+
|
|
47
|
+
with contextlib.redirect_stderr(captured):
|
|
48
|
+
for name, strict in _STRATEGIES:
|
|
49
|
+
try:
|
|
50
|
+
exported_program = torch.export.export(
|
|
51
|
+
model, example_inputs, dynamic_shapes=dynamic_shapes, strict=strict
|
|
52
|
+
)
|
|
53
|
+
except Exception as exc: # noqa: BLE001 - a failed strategy means try the next
|
|
54
|
+
last_exception = exc
|
|
55
|
+
else:
|
|
56
|
+
strategy_used = name
|
|
57
|
+
break
|
|
58
|
+
|
|
59
|
+
if exported_program is None:
|
|
60
|
+
return CaptureResult(
|
|
61
|
+
success=False,
|
|
62
|
+
capture_strategy=None,
|
|
63
|
+
exception=last_exception,
|
|
64
|
+
stderr=captured.getvalue(),
|
|
65
|
+
)
|
|
66
|
+
|
|
67
|
+
try:
|
|
68
|
+
onnx_program = torch.onnx.export(exported_program, verbose=False, report=False)
|
|
69
|
+
except Exception as exc: # noqa: BLE001 - a failed translation is a failed capture
|
|
70
|
+
return CaptureResult(
|
|
71
|
+
success=False,
|
|
72
|
+
capture_strategy=strategy_used,
|
|
73
|
+
exception=exc,
|
|
74
|
+
stderr=captured.getvalue(),
|
|
75
|
+
)
|
|
76
|
+
|
|
77
|
+
if onnx_program is None:
|
|
78
|
+
# The type stub allows None for the legacy path; with an ExportedProgram it never is.
|
|
79
|
+
return CaptureResult(
|
|
80
|
+
success=False,
|
|
81
|
+
capture_strategy=strategy_used,
|
|
82
|
+
exception=RuntimeError("torch.onnx.export returned None"),
|
|
83
|
+
)
|
|
84
|
+
|
|
85
|
+
proto = onnx_program.model_proto
|
|
86
|
+
return CaptureResult(
|
|
87
|
+
success=True,
|
|
88
|
+
capture_strategy=strategy_used,
|
|
89
|
+
onnx_program=onnx_program,
|
|
90
|
+
opset=proto.opset_import[0].version if proto.opset_import else None,
|
|
91
|
+
op_types=[node.op_type for node in proto.graph.node],
|
|
92
|
+
stderr=captured.getvalue(),
|
|
93
|
+
)
|