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/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()
@@ -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
+ )