modelport-cli 0.1.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.
modelport/__init__.py ADDED
@@ -0,0 +1,3 @@
1
+ """ModelPort: prepare models for Flutter apps."""
2
+
3
+ __version__ = "0.1.0"
modelport/bundle.py ADDED
@@ -0,0 +1,83 @@
1
+ """A bundle folder: modelport.json plus the files it points to."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from pathlib import Path
6
+ from typing import Any
7
+
8
+ from .hashing import sha256_file
9
+ from .manifest import MANIFEST_FILENAME, FileRef, Manifest
10
+
11
+ _SAFE_CHECK_SHA = "0" * 64
12
+
13
+
14
+ class Bundle:
15
+ """Write and check the files of one model bundle."""
16
+
17
+ def __init__(self, root: str | Path) -> None:
18
+ self.root = Path(root)
19
+
20
+ def path(self, relative: str) -> Path:
21
+ """Absolute path for a bundle-relative path. Rejects unsafe paths."""
22
+ FileRef(path=relative, size=1, sha256=_SAFE_CHECK_SHA) # validates the path rules
23
+ return self.root / relative
24
+
25
+ def add(self, relative: str) -> FileRef:
26
+ """Reference a file that is already inside the bundle."""
27
+ target = self.path(relative)
28
+ if not target.is_file():
29
+ raise FileNotFoundError(f"{target} does not exist")
30
+ return FileRef(path=relative, size=target.stat().st_size, sha256=sha256_file(target))
31
+
32
+ def write_bytes(self, relative: str, data: bytes) -> FileRef:
33
+ target = self.path(relative)
34
+ target.parent.mkdir(parents=True, exist_ok=True)
35
+ target.write_bytes(data)
36
+ return self.add(relative)
37
+
38
+ def write_text(self, relative: str, text: str) -> FileRef:
39
+ return self.write_bytes(relative, text.encode("utf-8"))
40
+
41
+ @property
42
+ def manifest_path(self) -> Path:
43
+ return self.root / MANIFEST_FILENAME
44
+
45
+ def write_manifest(self, manifest: Manifest) -> Path:
46
+ self.root.mkdir(parents=True, exist_ok=True)
47
+ self.manifest_path.write_text(manifest.to_json(), encoding="utf-8")
48
+ return self.manifest_path
49
+
50
+ def read_manifest(self) -> Manifest:
51
+ return Manifest.from_file(self.manifest_path)
52
+
53
+ def problems(self, manifest: Manifest) -> list[str]:
54
+ """Bundle files that are missing or do not match their size or sha256."""
55
+ found: list[str] = []
56
+ for ref in manifest.files():
57
+ if ref.path is None:
58
+ continue
59
+ target = self.root / ref.path
60
+ if not target.is_file():
61
+ found.append(f"{ref.path}: missing")
62
+ elif target.stat().st_size != ref.size:
63
+ found.append(
64
+ f"{ref.path}: size is {target.stat().st_size}, manifest says {ref.size}"
65
+ )
66
+ elif sha256_file(target) != ref.sha256:
67
+ found.append(f"{ref.path}: sha256 does not match")
68
+ return found
69
+
70
+ def refresh(self, manifest: Manifest) -> Manifest:
71
+ """Recompute size and sha256 of every bundle file the manifest points to."""
72
+ data = manifest.model_dump(by_alias=True, exclude_none=True, mode="json")
73
+ return Manifest.model_validate(self._refresh_refs(data))
74
+
75
+ def _refresh_refs(self, node: Any) -> Any:
76
+ if isinstance(node, dict):
77
+ if "path" in node and "sha256" in node and "size" in node:
78
+ ref = self.add(node["path"])
79
+ return {**node, "size": ref.size, "sha256": ref.sha256}
80
+ return {key: self._refresh_refs(value) for key, value in node.items()}
81
+ if isinstance(node, list):
82
+ return [self._refresh_refs(item) for item in node]
83
+ return node
modelport/cli.py ADDED
@@ -0,0 +1,527 @@
1
+ """The `modelport` command line tool."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+ from dataclasses import asdict
7
+ from pathlib import Path
8
+ from typing import Annotated
9
+
10
+ import typer
11
+ from pydantic import ValidationError
12
+ from rich.console import Console
13
+ from rich.table import Table
14
+
15
+ from . import __version__
16
+ from .doctor import collect_report
17
+ from .errors import ModelPortError
18
+ from .inspection import ModelInfo, TensorInfo, inspect_model
19
+ from .manifest import MANIFEST_FILENAME, Manifest
20
+ from .manifest.schema import render_schema
21
+
22
+ app = typer.Typer(
23
+ name="modelport",
24
+ help="Prepare PyTorch, Hugging Face, and GGUF models for Flutter apps.",
25
+ no_args_is_help=True,
26
+ add_completion=False,
27
+ pretty_exceptions_show_locals=False,
28
+ )
29
+ console = Console()
30
+ err_console = Console(stderr=True)
31
+
32
+
33
+ def _print_version(value: bool) -> None:
34
+ if value:
35
+ console.print(f"modelport {__version__}")
36
+ raise typer.Exit()
37
+
38
+
39
+ @app.callback()
40
+ def main(
41
+ version: Annotated[
42
+ bool,
43
+ typer.Option(
44
+ "--version",
45
+ "-V",
46
+ callback=_print_version,
47
+ is_eager=True,
48
+ help="Show the version and exit.",
49
+ ),
50
+ ] = False,
51
+ ) -> None:
52
+ """Prepare PyTorch, Hugging Face, and GGUF models for Flutter apps."""
53
+
54
+
55
+ @app.command()
56
+ def schema(
57
+ output: Annotated[
58
+ Path | None,
59
+ typer.Option(
60
+ "--output",
61
+ "-o",
62
+ help="File or folder to write manifest.schema.json to. Prints to stdout if omitted.",
63
+ ),
64
+ ] = None,
65
+ ) -> None:
66
+ """Print or write the JSON Schema for modelport.json."""
67
+ text = render_schema()
68
+ if output is None:
69
+ typer.echo(text, nl=False)
70
+ return
71
+ target = output / "manifest.schema.json" if output.is_dir() else output
72
+ target.write_text(text, encoding="utf-8")
73
+ err_console.print(f"Wrote {target}")
74
+
75
+
76
+ @app.command()
77
+ def validate(
78
+ paths: Annotated[
79
+ list[Path],
80
+ typer.Argument(help="modelport.json files, or bundle folders that contain one."),
81
+ ],
82
+ ) -> None:
83
+ """Check that manifests follow the spec. Exits with code 1 if any is invalid."""
84
+ failed = 0
85
+ for path in paths:
86
+ file = path / MANIFEST_FILENAME if path.is_dir() else path
87
+ try:
88
+ manifest = Manifest.from_file(file)
89
+ except ValidationError as error:
90
+ failed += 1
91
+ console.print(f"[red]✗[/red] {file}")
92
+ for issue in error.errors():
93
+ where = ".".join(str(part) for part in issue["loc"]) or "(root)"
94
+ console.print(f" {where}: {issue['msg']}", markup=False, highlight=False)
95
+ except (OSError, ValueError) as error:
96
+ failed += 1
97
+ console.print(f"[red]✗[/red] {file}")
98
+ console.print(f" {error}", markup=False, highlight=False)
99
+ else:
100
+ console.print(
101
+ f"[green]✓[/green] {file} {manifest.id} {manifest.version}, "
102
+ f"{manifest.task}, {len(manifest.variants)} variant(s)"
103
+ )
104
+ if failed:
105
+ raise typer.Exit(code=1)
106
+
107
+
108
+ @app.command()
109
+ def doctor() -> None:
110
+ """Check Python, optional packages, and disk space."""
111
+ report = collect_report()
112
+ summary = Table.grid(padding=(0, 2))
113
+ summary.add_row("modelport", report.modelport_version)
114
+ summary.add_row("python", report.python_version)
115
+ summary.add_row("platform", report.platform)
116
+ summary.add_row("free disk", f"{report.free_disk_bytes / 1000**3:.1f} GB")
117
+ console.print(summary)
118
+ console.print()
119
+
120
+ table = Table(title="Optional packages", title_justify="left", show_edge=False)
121
+ table.add_column("package")
122
+ table.add_column("extra")
123
+ table.add_column("version")
124
+ for package in report.packages:
125
+ version = package.version or "[yellow]not installed[/yellow]"
126
+ table.add_row(package.name, package.extra, version)
127
+ console.print(table)
128
+
129
+ for problem in report.problems:
130
+ console.print(f"[red]![/red] {problem}", highlight=False)
131
+ if report.hints:
132
+ console.print("\nTo enable more formats:")
133
+ for hint in report.hints:
134
+ console.print(f" {hint}", markup=False, highlight=False)
135
+ if not report.problems and not report.hints:
136
+ console.print("\n[green]Everything is installed.[/green]")
137
+
138
+
139
+ @app.command()
140
+ def export(
141
+ source: Annotated[
142
+ str,
143
+ typer.Argument(
144
+ help="Model to export: torchvision:<name>, hf:<repo or folder>, or file:<script.py>."
145
+ ),
146
+ ],
147
+ target: Annotated[
148
+ list[str],
149
+ typer.Option(
150
+ "--target", "-t", help="Format to export: onnx, executorch. Repeat or comma-separate."
151
+ ),
152
+ ] = ["onnx"], # noqa: B006 - Typer reads list defaults
153
+ out: Annotated[Path, typer.Option("--out", "-o", help="Folder for bundles.")] = Path("dist"),
154
+ license: Annotated[
155
+ str | None, typer.Option(help="SPDX license of the weights, if the source has none.")
156
+ ] = None,
157
+ version: Annotated[str, typer.Option(help="Version of this bundle.")] = "1.0.0",
158
+ sample_image: Annotated[
159
+ Path | None,
160
+ typer.Option(
161
+ "--sample-image", help="Picture used for golden data. A fixed one by default."
162
+ ),
163
+ ] = None,
164
+ force: Annotated[bool, typer.Option("--force", help="Replace an existing bundle.")] = False,
165
+ image_size: Annotated[
166
+ int | None,
167
+ typer.Option("--image-size", help="Square input size for detectors that accept any size."),
168
+ ] = None,
169
+ ) -> None:
170
+ """Convert a model and write a bundle with modelport.json and golden test data."""
171
+ from .pipeline import export_bundle
172
+ from .preprocess import load_image
173
+ from .sources import load_source
174
+
175
+ targets = [t.strip() for item in target for t in item.split(",") if t.strip()]
176
+ try:
177
+ with console.status(f"Loading {source}"):
178
+ model = load_source(source, license=license, image_size=image_size)
179
+ image = load_image(sample_image) if sample_image else None
180
+ with console.status(f"Exporting {model.id} to {', '.join(targets)}"):
181
+ bundle, manifest = export_bundle(
182
+ model, out, targets, image=image, version=version, overwrite=force
183
+ )
184
+ except ModelPortError as error:
185
+ err_console.print(f"[red]Error:[/red] {error}", highlight=False)
186
+ raise typer.Exit(code=1) from error
187
+
188
+ table = Table(title=f"Wrote {bundle.root}", title_justify="left", show_edge=False)
189
+ table.add_column("file")
190
+ table.add_column("size", justify="right")
191
+ for ref in manifest.files():
192
+ if ref.path is not None:
193
+ table.add_row(ref.path, _format_size(ref.size))
194
+ table.add_row("modelport.json", _format_size(bundle.manifest_path.stat().st_size))
195
+ console.print(table)
196
+ console.print(f"\nNext: modelport verify {bundle.root}", highlight=False)
197
+
198
+
199
+ @app.command()
200
+ def verify(
201
+ bundle: Annotated[Path, typer.Argument(help="Bundle folder that contains modelport.json.")],
202
+ ) -> None:
203
+ """Run every variant on the golden input and compare with the expected output."""
204
+ from .verify import verify_bundle
205
+
206
+ try:
207
+ with console.status(f"Verifying {bundle}"):
208
+ results = verify_bundle(bundle)
209
+ except ModelPortError as error:
210
+ err_console.print(f"[red]Error:[/red] {error}", highlight=False)
211
+ raise typer.Exit(code=1) from error
212
+
213
+ table = Table(show_edge=False)
214
+ table.add_column("variant", no_wrap=True)
215
+ for column in ("output", "max diff", "cosine", "top-1", "result"):
216
+ table.add_column(column, no_wrap=True)
217
+ for result in results:
218
+ if result.skipped:
219
+ table.add_row(result.variant_id, "", "", "", "", "[yellow]skipped[/]")
220
+ continue
221
+ for check in result.outputs:
222
+ ok = check.within_tolerance and check.top1_match is not False
223
+ top1 = {None: "", True: "same", False: "[red]different[/]"}[check.top1_match]
224
+ table.add_row(
225
+ result.variant_id,
226
+ check.name,
227
+ f"{check.max_abs_diff:.2e}",
228
+ f"{check.cosine:.6f}",
229
+ top1,
230
+ "[green]pass[/]" if ok else "[red]fail[/]",
231
+ )
232
+ console.print(table)
233
+ for result in results:
234
+ if result.skipped:
235
+ console.print(f"[yellow]![/] {result.variant_id}: {result.skipped}", highlight=False)
236
+
237
+ failed = [r for r in results if not r.passed and r.skipped is None]
238
+ if failed or not any(r.passed for r in results):
239
+ raise typer.Exit(code=1)
240
+
241
+
242
+ @app.command()
243
+ def quantize(
244
+ bundle: Annotated[Path, typer.Argument(help="Bundle folder with an onnx fp32 variant.")],
245
+ fp16: Annotated[
246
+ bool, typer.Option("--fp16", help="Add an fp16 variant (about half size).")
247
+ ] = False,
248
+ int8: Annotated[
249
+ bool, typer.Option("--int8", help="Add an int8 variant of MatMul and Gemm weights.")
250
+ ] = False,
251
+ allow_top1_change: Annotated[
252
+ bool, typer.Option(help="Keep a variant even if it changes the golden top-1 class.")
253
+ ] = False,
254
+ ) -> None:
255
+ """Add smaller ONNX variants and record how far they drift from the original."""
256
+ from .quantize import Kind, quantize_bundle
257
+
258
+ kinds: list[Kind] = []
259
+ if fp16:
260
+ kinds.append("fp16")
261
+ if int8:
262
+ kinds.append("int8")
263
+ if not kinds:
264
+ err_console.print("Choose at least one of --fp16 or --int8.")
265
+ raise typer.Exit(code=2)
266
+ try:
267
+ with console.status(f"Quantizing {bundle}"):
268
+ results = quantize_bundle(bundle, kinds, allow_top1_change=allow_top1_change)
269
+ except ModelPortError as error:
270
+ err_console.print(f"[red]Error:[/red] {error}", highlight=False)
271
+ raise typer.Exit(code=1) from error
272
+
273
+ table = Table(show_edge=False)
274
+ for column in ("variant", "size", "of fp32", "max |diff|", "tolerance", "top-1"):
275
+ table.add_column(column)
276
+ for result in results:
277
+ variant = result.variant
278
+ top1 = {None: "", True: "same", False: "[red]different[/]"}[result.top1_match]
279
+ tolerance = variant.tolerance.atol if variant.tolerance else ""
280
+ table.add_row(
281
+ variant.id,
282
+ _format_size(variant.file.size),
283
+ f"{result.size_ratio:.0%}",
284
+ f"{result.max_abs_diff:.2e}",
285
+ str(tolerance),
286
+ top1,
287
+ )
288
+ console.print(table)
289
+ console.print(f"\nNext: modelport verify {bundle}", highlight=False)
290
+
291
+
292
+ @app.command()
293
+ def pack(
294
+ bundle: Annotated[Path, typer.Argument(help="Bundle folder that contains modelport.json.")],
295
+ ) -> None:
296
+ """Refresh file sizes and hashes after manual edits, and list unlisted files."""
297
+ from .pack import pack_bundle
298
+
299
+ try:
300
+ result = pack_bundle(bundle)
301
+ except ModelPortError as error:
302
+ err_console.print(f"[red]Error:[/red] {error}", highlight=False)
303
+ raise typer.Exit(code=1) from error
304
+ for path in result.changed:
305
+ console.print(f"updated hash: {path}", highlight=False)
306
+ for path in result.stray:
307
+ console.print(
308
+ f"[yellow]![/] not in manifest, will not be published: {path}", highlight=False
309
+ )
310
+ console.print(
311
+ f"[green]✓[/green] {result.manifest.id} {result.manifest.version}, "
312
+ f"{_format_size(result.total_bytes)} in {len(result.manifest.files())} files"
313
+ )
314
+
315
+
316
+ @app.command()
317
+ def publish(
318
+ bundle: Annotated[Path, typer.Argument(help="Bundle folder that contains modelport.json.")],
319
+ hf: Annotated[
320
+ str | None, typer.Option("--hf", help="Hugging Face repo id, like org/name.")
321
+ ] = None,
322
+ github: Annotated[
323
+ str | None,
324
+ typer.Option("--github", help="GitHub repo, like owner/name. Uploads release assets."),
325
+ ] = None,
326
+ tag: Annotated[str, typer.Option(help="Release tag for --github.")] = "models",
327
+ private: Annotated[bool, typer.Option(help="Create the Hugging Face repo as private.")] = False,
328
+ ) -> None:
329
+ """Upload a bundle to the Hugging Face Hub or to a GitHub release.
330
+
331
+ For Hugging Face, log in first with `hf auth login`. For GitHub, log in with `gh auth login`.
332
+ """
333
+ from .publish import publish_bundle, publish_to_github
334
+
335
+ if (hf is None) == (github is None):
336
+ err_console.print("Choose exactly one of --hf or --github.")
337
+ raise typer.Exit(code=2)
338
+ try:
339
+ with console.status(f"Uploading {bundle}"):
340
+ if github is not None:
341
+ result = publish_to_github(bundle, github, tag)
342
+ else:
343
+ assert hf is not None
344
+ result = publish_bundle(bundle, hf, private=private)
345
+ except ModelPortError as error:
346
+ err_console.print(f"[red]Error:[/red] {error}", highlight=False)
347
+ raise typer.Exit(code=1) from error
348
+ console.print(f"[green]✓[/green] Uploaded {len(result.files)} files to {result.url}")
349
+ console.print(f"Load it in Flutter with: {result.location}", highlight=False)
350
+
351
+
352
+ @app.command("import-gguf")
353
+ def import_gguf(
354
+ source: Annotated[
355
+ str,
356
+ typer.Argument(
357
+ help="Hugging Face repo like Qwen/Qwen2.5-0.5B-Instruct-GGUF, or a .gguf file."
358
+ ),
359
+ ],
360
+ quant: Annotated[
361
+ list[str],
362
+ typer.Option("--quant", "-q", help="Quantizations to include, in order of preference."),
363
+ ] = ["q4_k_m"], # noqa: B006 - Typer reads list defaults
364
+ context: Annotated[int, typer.Option(help="Context window to allocate on device.")] = 4096,
365
+ license: Annotated[str | None, typer.Option(help="SPDX license, if the repo has none.")] = None,
366
+ model_id: Annotated[
367
+ str | None, typer.Option("--id", help="Bundle id. Taken from the source by default.")
368
+ ] = None,
369
+ revision: Annotated[str | None, typer.Option(help="Repo branch, tag, or commit.")] = None,
370
+ out: Annotated[Path, typer.Option("--out", "-o", help="Folder for bundles.")] = Path("dist"),
371
+ force: Annotated[bool, typer.Option("--force", help="Replace an existing bundle.")] = False,
372
+ ) -> None:
373
+ """Describe a ready-made GGUF language model with a manifest.
374
+
375
+ Hub files are not downloaded: the manifest points at them with pinned URLs.
376
+ """
377
+ import shutil
378
+
379
+ from .bundle import Bundle
380
+ from .gguf_import import import_from_file, import_from_hub
381
+
382
+ quants = [q.strip() for item in quant for q in item.split(",") if q.strip()]
383
+ local = Path(source)
384
+ try:
385
+ if local.is_file():
386
+ if license is None:
387
+ raise ModelPortError("local GGUF files need --license")
388
+ shutil.rmtree(out / ".import", ignore_errors=True)
389
+ staging = Bundle(out / ".import")
390
+ result = import_from_file(
391
+ local, staging, context=context, license=license, model_id=model_id
392
+ )
393
+ root = out / result.manifest.id
394
+ else:
395
+ repo = source.removeprefix("hf:")
396
+ with console.status(f"Reading {repo} from the Hugging Face Hub"):
397
+ result = import_from_hub(
398
+ repo,
399
+ quants,
400
+ revision=revision,
401
+ context=context,
402
+ license=license,
403
+ model_id=model_id,
404
+ )
405
+ root = out / result.manifest.id
406
+ staging = None
407
+ if root.exists() and any(root.iterdir()):
408
+ if not force:
409
+ raise ModelPortError(f"{root} already exists. Use --force to replace it.")
410
+ shutil.rmtree(root)
411
+ if staging is not None:
412
+ root.parent.mkdir(parents=True, exist_ok=True)
413
+ staging.root.rename(root)
414
+ bundle = Bundle(root)
415
+ bundle.write_manifest(result.manifest)
416
+ except ModelPortError as error:
417
+ err_console.print(f"[red]Error:[/red] {error}", highlight=False)
418
+ raise typer.Exit(code=1) from error
419
+
420
+ console.print(f"Wrote {bundle.manifest_path}", highlight=False)
421
+ table = Table(show_edge=False)
422
+ table.add_column("variant", no_wrap=True)
423
+ table.add_column("size", justify="right")
424
+ table.add_column("min RAM", justify="right")
425
+ for variant in result.manifest.variants:
426
+ table.add_row(variant.id, _format_size(variant.file.size), f"{variant.min_ram_mb} MB")
427
+ console.print(table)
428
+ for warning in result.warnings:
429
+ console.print(f"[yellow]![/] {warning}", highlight=False)
430
+ console.print(
431
+ f"\nNext: modelport publish {bundle.root} --hf <your-org>/<name>", highlight=False
432
+ )
433
+
434
+
435
+ @app.command("gen-dart")
436
+ def gen_dart(
437
+ manifest_path: Annotated[
438
+ Path, typer.Argument(help="modelport.json, or a bundle folder that contains one.")
439
+ ],
440
+ out: Annotated[Path, typer.Option("--out", "-o", help="Folder for the .dart file.")] = Path(
441
+ "lib/models"
442
+ ),
443
+ location: Annotated[
444
+ str | None,
445
+ typer.Option(help="Default location for load(), such as hf://org/name."),
446
+ ] = None,
447
+ ) -> None:
448
+ """Generate a typed Dart wrapper with named inputs and outputs."""
449
+ from .codegen.dart import file_name, format_with_dart, generate_dart
450
+ from .manifest import MANIFEST_FILENAME, Manifest
451
+
452
+ path = manifest_path / MANIFEST_FILENAME if manifest_path.is_dir() else manifest_path
453
+ try:
454
+ manifest = Manifest.from_file(path)
455
+ source = generate_dart(manifest, location=location)
456
+ formatted = format_with_dart(source)
457
+ except (ModelPortError, ValueError, OSError) as error:
458
+ err_console.print(f"[red]Error:[/red] {error}", highlight=False)
459
+ raise typer.Exit(code=1) from error
460
+ out.mkdir(parents=True, exist_ok=True)
461
+ target = out / file_name(manifest.id)
462
+ target.write_text(formatted or source, encoding="utf-8")
463
+ console.print(f"Wrote {target}", highlight=False)
464
+ if formatted is None:
465
+ console.print("Dart is not installed, so the file was not formatted.", highlight=False)
466
+
467
+
468
+ @app.command("inspect")
469
+ def inspect_command(
470
+ path: Annotated[Path, typer.Argument(help="A .onnx, .pte, or .gguf model file.")],
471
+ as_json: Annotated[bool, typer.Option("--json", help="Print machine-readable JSON.")] = False,
472
+ ) -> None:
473
+ """Show a model file's inputs, outputs, and metadata."""
474
+ try:
475
+ info = inspect_model(path)
476
+ except ModelPortError as error:
477
+ err_console.print(f"[red]Error:[/red] {error}", highlight=False)
478
+ raise typer.Exit(code=1) from error
479
+
480
+ if as_json:
481
+ data = asdict(info)
482
+ data["path"] = str(info.path)
483
+ typer.echo(json.dumps(data, indent=2))
484
+ return
485
+ _print_model_info(info)
486
+
487
+
488
+ def _format_size(size: int) -> str:
489
+ value = float(size)
490
+ for unit in ("B", "KB", "MB", "GB"):
491
+ if value < 1000 or unit == "GB":
492
+ return f"{value:.1f} {unit}" if unit != "B" else f"{size} B"
493
+ value /= 1000
494
+ raise AssertionError("unreachable")
495
+
496
+
497
+ def _tensor_table(title: str, tensors: list[TensorInfo]) -> Table:
498
+ table = Table(title=title, title_justify="left", show_edge=False)
499
+ table.add_column("name")
500
+ table.add_column("dtype")
501
+ table.add_column("shape")
502
+ for tensor in tensors:
503
+ table.add_row(tensor.name, tensor.dtype, str(tensor.shape))
504
+ return table
505
+
506
+
507
+ def _print_model_info(info: ModelInfo) -> None:
508
+ summary = Table.grid(padding=(0, 2))
509
+ summary.add_row("file", str(info.path))
510
+ summary.add_row("format", info.format)
511
+ summary.add_row("size", f"{_format_size(info.size)} ({info.size:,} bytes)")
512
+ summary.add_row("sha256", info.sha256)
513
+ console.print(summary)
514
+ if info.inputs:
515
+ console.print()
516
+ console.print(_tensor_table("Inputs", info.inputs))
517
+ if info.outputs:
518
+ console.print()
519
+ console.print(_tensor_table("Outputs", info.outputs))
520
+ if info.metadata:
521
+ console.print()
522
+ meta = Table(title="Metadata", title_justify="left", show_edge=False)
523
+ meta.add_column("key")
524
+ meta.add_column("value")
525
+ for key, value in info.metadata.items():
526
+ meta.add_row(key, value)
527
+ console.print(meta)
@@ -0,0 +1 @@
1
+ """Code generation from manifests."""