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 +3 -0
- modelport/bundle.py +83 -0
- modelport/cli.py +527 -0
- modelport/codegen/__init__.py +1 -0
- modelport/codegen/dart.py +227 -0
- modelport/doctor.py +88 -0
- modelport/errors.py +18 -0
- modelport/exporters/__init__.py +55 -0
- modelport/exporters/executorch.py +46 -0
- modelport/exporters/onnx.py +63 -0
- modelport/gguf_import.py +197 -0
- modelport/golden.py +108 -0
- modelport/hashing.py +17 -0
- modelport/inspection/__init__.py +87 -0
- modelport/inspection/executorch.py +32 -0
- modelport/inspection/gguf.py +42 -0
- modelport/inspection/onnx.py +44 -0
- modelport/manifest/__init__.py +39 -0
- modelport/manifest/base.py +22 -0
- modelport/manifest/files.py +59 -0
- modelport/manifest/models.py +241 -0
- modelport/manifest/postprocess.py +79 -0
- modelport/manifest/schema.py +49 -0
- modelport/manifest/tensors.py +158 -0
- modelport/pack.py +48 -0
- modelport/pipeline.py +109 -0
- modelport/preprocess.py +121 -0
- modelport/publish.py +191 -0
- modelport/py.typed +0 -0
- modelport/quantize.py +120 -0
- modelport/runtimes.py +52 -0
- modelport/sources/__init__.py +85 -0
- modelport/sources/file.py +46 -0
- modelport/sources/hf.py +265 -0
- modelport/sources/torchvision.py +78 -0
- modelport/verify.py +92 -0
- modelport_cli-0.1.0.dist-info/METADATA +152 -0
- modelport_cli-0.1.0.dist-info/RECORD +41 -0
- modelport_cli-0.1.0.dist-info/WHEEL +4 -0
- modelport_cli-0.1.0.dist-info/entry_points.txt +2 -0
- modelport_cli-0.1.0.dist-info/licenses/LICENSE +202 -0
modelport/preprocess.py
ADDED
|
@@ -0,0 +1,121 @@
|
|
|
1
|
+
"""Reference implementation of manifest image preprocessing.
|
|
2
|
+
|
|
3
|
+
The Dart packages implement the same steps, and cross-language tests compare the two.
|
|
4
|
+
See docs/spec.md, "Image preprocessing", for the exact rules.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
from pathlib import Path
|
|
10
|
+
|
|
11
|
+
import numpy as np
|
|
12
|
+
from PIL import Image, ImageOps
|
|
13
|
+
|
|
14
|
+
from .manifest.tensors import DType, ImagePreprocess, InputSpec, ResizeSpec
|
|
15
|
+
|
|
16
|
+
_PIL_FILTERS = {
|
|
17
|
+
"bilinear": Image.Resampling.BILINEAR,
|
|
18
|
+
"bicubic": Image.Resampling.BICUBIC,
|
|
19
|
+
"nearest": Image.Resampling.NEAREST,
|
|
20
|
+
}
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def load_image(path: str | Path) -> Image.Image:
|
|
24
|
+
"""Open an image, apply its EXIF orientation, and convert it to 8-bit RGB."""
|
|
25
|
+
with Image.open(path) as image:
|
|
26
|
+
return ImageOps.exif_transpose(image).convert("RGB")
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
def resized_size(width: int, height: int, spec: ResizeSpec) -> tuple[int, int]:
|
|
30
|
+
"""Output (width, height) for a resize rule."""
|
|
31
|
+
if spec.size is not None:
|
|
32
|
+
out_height, out_width = spec.size
|
|
33
|
+
return out_width, out_height
|
|
34
|
+
assert spec.shorter_side is not None
|
|
35
|
+
short = spec.shorter_side
|
|
36
|
+
if width <= height:
|
|
37
|
+
return short, int(short * height / width)
|
|
38
|
+
return int(short * width / height), short
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def resize(image: Image.Image, spec: ResizeSpec) -> np.ndarray:
|
|
42
|
+
"""Resize to a float32 array of shape (height, width, 3) with values in 0-255."""
|
|
43
|
+
width, height = resized_size(image.width, image.height, spec)
|
|
44
|
+
if spec.antialias or spec.method == "nearest":
|
|
45
|
+
if (width, height) != image.size:
|
|
46
|
+
image = image.resize((width, height), resample=_PIL_FILTERS[spec.method])
|
|
47
|
+
return np.asarray(image, dtype=np.float32)
|
|
48
|
+
return bilinear_half_pixel(np.asarray(image, dtype=np.float32), height, width)
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def _axis(in_size: int, out_size: int) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
|
|
52
|
+
scale = in_size / out_size
|
|
53
|
+
src = np.maximum((np.arange(out_size, dtype=np.float64) + 0.5) * scale - 0.5, 0.0)
|
|
54
|
+
low = np.minimum(np.floor(src).astype(np.int64), in_size - 1)
|
|
55
|
+
high = np.minimum(low + 1, in_size - 1)
|
|
56
|
+
return low, high, src - low
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def bilinear_half_pixel(array: np.ndarray, out_height: int, out_width: int) -> np.ndarray:
|
|
60
|
+
"""Bilinear resize with half-pixel centers and no antialiasing, in floating point.
|
|
61
|
+
|
|
62
|
+
Matches PyTorch `interpolate(mode="bilinear", align_corners=False, antialias=False)`
|
|
63
|
+
and OpenCV `INTER_LINEAR`.
|
|
64
|
+
"""
|
|
65
|
+
data = array.astype(np.float64)
|
|
66
|
+
y_low, y_high, y_frac = _axis(data.shape[0], out_height)
|
|
67
|
+
x_low, x_high, x_frac = _axis(data.shape[1], out_width)
|
|
68
|
+
y_frac = y_frac[:, None, None]
|
|
69
|
+
rows = data[y_low] * (1 - y_frac) + data[y_high] * y_frac
|
|
70
|
+
x_frac = x_frac[None, :, None]
|
|
71
|
+
out = rows[:, x_low] * (1 - x_frac) + rows[:, x_high] * x_frac
|
|
72
|
+
return out.astype(np.float32)
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def center_crop(array: np.ndarray, crop_height: int, crop_width: int) -> np.ndarray:
|
|
76
|
+
"""Crop the center, padding with zeros first if the image is too small."""
|
|
77
|
+
height, width = array.shape[:2]
|
|
78
|
+
if crop_height > height or crop_width > width:
|
|
79
|
+
dh, dw = max(crop_height - height, 0), max(crop_width - width, 0)
|
|
80
|
+
array = np.pad(array, ((dh // 2, (dh + 1) // 2), (dw // 2, (dw + 1) // 2), (0, 0)))
|
|
81
|
+
height, width = array.shape[:2]
|
|
82
|
+
# Python's round() is round-half-to-even, as in torchvision.
|
|
83
|
+
top = round((height - crop_height) / 2)
|
|
84
|
+
left = round((width - crop_width) / 2)
|
|
85
|
+
return array[top : top + crop_height, left : left + crop_width]
|
|
86
|
+
|
|
87
|
+
|
|
88
|
+
def preprocess_array(image: Image.Image, pre: ImagePreprocess) -> np.ndarray:
|
|
89
|
+
"""Run resize, crop, channel order, scale, and normalize. Returns (H, W, 3) float64."""
|
|
90
|
+
array = resize(image.convert("RGB"), pre.resize)
|
|
91
|
+
if pre.center_crop is not None:
|
|
92
|
+
array = center_crop(array, *pre.center_crop)
|
|
93
|
+
if pre.color == "BGR":
|
|
94
|
+
array = array[..., ::-1]
|
|
95
|
+
values = array.astype(np.float64) * pre.scale
|
|
96
|
+
return (values - np.asarray(pre.mean)) / np.asarray(pre.std)
|
|
97
|
+
|
|
98
|
+
|
|
99
|
+
def preprocess_image(image: Image.Image, spec: InputSpec) -> np.ndarray:
|
|
100
|
+
"""Build the input tensor for one image, including the batch dimension."""
|
|
101
|
+
if spec.preprocess is None:
|
|
102
|
+
raise ValueError(f"input '{spec.name}' has no image preprocessing")
|
|
103
|
+
values = preprocess_array(image, spec.preprocess)
|
|
104
|
+
if spec.dtype is DType.UINT8:
|
|
105
|
+
tensor = np.clip(np.rint(values), 0, 255).astype(np.uint8)
|
|
106
|
+
elif spec.dtype is DType.FLOAT16:
|
|
107
|
+
tensor = values.astype(np.float16)
|
|
108
|
+
else:
|
|
109
|
+
tensor = values.astype(np.float32)
|
|
110
|
+
tensor = tensor.transpose(2, 0, 1)[None] if spec.layout == "NCHW" else tensor[None]
|
|
111
|
+
for axis, (expected, actual) in enumerate(zip(spec.shape, tensor.shape, strict=True)):
|
|
112
|
+
if expected != -1 and expected != actual:
|
|
113
|
+
raise ValueError(
|
|
114
|
+
f"input '{spec.name}' dimension {axis} is {actual}, manifest says {expected}"
|
|
115
|
+
)
|
|
116
|
+
return np.ascontiguousarray(tensor)
|
|
117
|
+
|
|
118
|
+
|
|
119
|
+
def to_le_bytes(array: np.ndarray) -> bytes:
|
|
120
|
+
"""Raw little-endian bytes, the format of golden files."""
|
|
121
|
+
return np.ascontiguousarray(array).astype(array.dtype.newbyteorder("<")).tobytes()
|
modelport/publish.py
ADDED
|
@@ -0,0 +1,191 @@
|
|
|
1
|
+
"""Upload a bundle to the Hugging Face Hub, where apps can load it with hf://."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import shutil
|
|
6
|
+
import subprocess
|
|
7
|
+
import tempfile
|
|
8
|
+
from collections.abc import Callable
|
|
9
|
+
from dataclasses import dataclass
|
|
10
|
+
from pathlib import Path
|
|
11
|
+
|
|
12
|
+
from .bundle import Bundle
|
|
13
|
+
from .errors import MissingDependencyError, ModelPortError
|
|
14
|
+
from .manifest import MANIFEST_FILENAME, Manifest
|
|
15
|
+
|
|
16
|
+
CARD_NAME = "README.md"
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
@dataclass(frozen=True)
|
|
20
|
+
class PublishResult:
|
|
21
|
+
repo_id: str
|
|
22
|
+
url: str
|
|
23
|
+
location: str
|
|
24
|
+
"""What an app passes to ModelPort.load()."""
|
|
25
|
+
files: list[str]
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
def model_card(manifest: Manifest, repo_id: str) -> str:
|
|
29
|
+
"""A short Hub model card describing the bundle and how to load it."""
|
|
30
|
+
rows = "\n".join(
|
|
31
|
+
f"| `{v.id}` | {v.runtime} | {v.precision} | {v.file.size / 1e6:.1f} MB |"
|
|
32
|
+
for v in manifest.variants
|
|
33
|
+
)
|
|
34
|
+
title = manifest.name or manifest.id
|
|
35
|
+
description = manifest.description or ""
|
|
36
|
+
return f"""---
|
|
37
|
+
license: {manifest.license.lower()}
|
|
38
|
+
tags:
|
|
39
|
+
- modelport
|
|
40
|
+
- flutter
|
|
41
|
+
- on-device
|
|
42
|
+
- {manifest.task}
|
|
43
|
+
---
|
|
44
|
+
|
|
45
|
+
# {title}
|
|
46
|
+
|
|
47
|
+
{description}
|
|
48
|
+
|
|
49
|
+
This is a [ModelPort](https://github.com/ayanparvaiz/modelport) bundle. The
|
|
50
|
+
`{MANIFEST_FILENAME}` file describes inputs, preprocessing, outputs, and checksums, so a
|
|
51
|
+
Flutter app can download and run it with one line.
|
|
52
|
+
|
|
53
|
+
| Variant | Runtime | Precision | Size |
|
|
54
|
+
|---|---|---|---|
|
|
55
|
+
{rows}
|
|
56
|
+
|
|
57
|
+
## Use in Flutter
|
|
58
|
+
|
|
59
|
+
```dart
|
|
60
|
+
final model = await ModelPort.load('hf://{repo_id}');
|
|
61
|
+
```
|
|
62
|
+
|
|
63
|
+
Source: `{manifest.source or "unknown"}`. License: {manifest.license}.
|
|
64
|
+
"""
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
def publish_bundle(path: str | Path, repo_id: str, *, private: bool = False) -> PublishResult:
|
|
68
|
+
"""Check the bundle, write a model card if missing, and upload the listed files."""
|
|
69
|
+
if repo_id.count("/") != 1:
|
|
70
|
+
raise ModelPortError(f"'{repo_id}' is not a repo id like 'org/name'")
|
|
71
|
+
bundle = Bundle(path)
|
|
72
|
+
manifest = bundle.read_manifest()
|
|
73
|
+
problems = bundle.problems(manifest)
|
|
74
|
+
if problems:
|
|
75
|
+
raise ModelPortError(
|
|
76
|
+
"bundle files do not match the manifest, run `modelport pack` first:\n "
|
|
77
|
+
+ "\n ".join(problems)
|
|
78
|
+
)
|
|
79
|
+
card = bundle.root / CARD_NAME
|
|
80
|
+
if not card.exists():
|
|
81
|
+
card.write_text(model_card(manifest, repo_id), encoding="utf-8")
|
|
82
|
+
|
|
83
|
+
# Only files inside the bundle are uploaded; https urls stay where they are.
|
|
84
|
+
files = [MANIFEST_FILENAME, CARD_NAME] + [ref.path for ref in manifest.files() if ref.path]
|
|
85
|
+
try:
|
|
86
|
+
from huggingface_hub import HfApi
|
|
87
|
+
except ImportError as error:
|
|
88
|
+
raise MissingDependencyError("Publishing to Hugging Face", "hf") from error
|
|
89
|
+
|
|
90
|
+
api = HfApi()
|
|
91
|
+
api.create_repo(repo_id, repo_type="model", private=private, exist_ok=True)
|
|
92
|
+
api.upload_folder(
|
|
93
|
+
folder_path=str(bundle.root),
|
|
94
|
+
repo_id=repo_id,
|
|
95
|
+
repo_type="model",
|
|
96
|
+
allow_patterns=files,
|
|
97
|
+
commit_message=f"modelport: {manifest.id} {manifest.version}",
|
|
98
|
+
)
|
|
99
|
+
return PublishResult(
|
|
100
|
+
repo_id=repo_id,
|
|
101
|
+
url=f"https://huggingface.co/{repo_id}",
|
|
102
|
+
location=f"hf://{repo_id}",
|
|
103
|
+
files=files,
|
|
104
|
+
)
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
def github_asset_name(model_id: str, path: str) -> str:
|
|
108
|
+
"""Release assets are flat, so 'onnx-fp32/model.onnx' becomes 'id--onnx-fp32--model.onnx'."""
|
|
109
|
+
return f"{model_id}--{path.replace('/', '--')}"
|
|
110
|
+
|
|
111
|
+
|
|
112
|
+
def github_manifest(manifest: Manifest, repo: str, tag: str) -> Manifest:
|
|
113
|
+
"""The manifest with every bundle path replaced by its release download URL."""
|
|
114
|
+
base = f"https://github.com/{repo}/releases/download/{tag}"
|
|
115
|
+
data = manifest.model_dump(by_alias=True, exclude_none=True, mode="json")
|
|
116
|
+
|
|
117
|
+
def rewrite(node: object) -> object:
|
|
118
|
+
if isinstance(node, dict):
|
|
119
|
+
if "path" in node and "sha256" in node and "size" in node:
|
|
120
|
+
name = github_asset_name(manifest.id, str(node["path"]))
|
|
121
|
+
return {"url": f"{base}/{name}", "size": node["size"], "sha256": node["sha256"]}
|
|
122
|
+
return {key: rewrite(value) for key, value in node.items()}
|
|
123
|
+
if isinstance(node, list):
|
|
124
|
+
return [rewrite(item) for item in node]
|
|
125
|
+
return node
|
|
126
|
+
|
|
127
|
+
return Manifest.model_validate(rewrite(data))
|
|
128
|
+
|
|
129
|
+
|
|
130
|
+
def publish_to_github(
|
|
131
|
+
path: str | Path,
|
|
132
|
+
repo: str,
|
|
133
|
+
tag: str,
|
|
134
|
+
*,
|
|
135
|
+
run: Callable[[list[str]], subprocess.CompletedProcess[str]] | None = None,
|
|
136
|
+
) -> PublishResult:
|
|
137
|
+
"""Upload a bundle as assets of a GitHub release, using the `gh` CLI.
|
|
138
|
+
|
|
139
|
+
The uploaded manifest, `<id>.json`, points at the other assets by URL, so apps
|
|
140
|
+
load the model from https://github.com/<repo>/releases/download/<tag>/<id>.json.
|
|
141
|
+
"""
|
|
142
|
+
if repo.count("/") != 1:
|
|
143
|
+
raise ModelPortError(f"'{repo}' is not a GitHub repo like 'owner/name'")
|
|
144
|
+
bundle = Bundle(path)
|
|
145
|
+
manifest = bundle.read_manifest()
|
|
146
|
+
problems = bundle.problems(manifest)
|
|
147
|
+
if problems:
|
|
148
|
+
raise ModelPortError(
|
|
149
|
+
"bundle files do not match the manifest, run `modelport pack` first:\n "
|
|
150
|
+
+ "\n ".join(problems)
|
|
151
|
+
)
|
|
152
|
+
execute = run or _run_gh
|
|
153
|
+
if shutil.which("gh") is None and run is None:
|
|
154
|
+
raise ModelPortError("publishing to GitHub needs the gh CLI: https://cli.github.com")
|
|
155
|
+
|
|
156
|
+
if execute(["gh", "release", "view", tag, "-R", repo]).returncode != 0:
|
|
157
|
+
created = execute(
|
|
158
|
+
[
|
|
159
|
+
"gh", "release", "create", tag, "-R", repo,
|
|
160
|
+
"--title", f"Model zoo {tag}",
|
|
161
|
+
"--notes", "ModelPort model bundles. Load them by their .json URL.",
|
|
162
|
+
"--latest=false",
|
|
163
|
+
]
|
|
164
|
+
) # fmt: skip
|
|
165
|
+
if created.returncode != 0:
|
|
166
|
+
raise ModelPortError(f"could not create release {tag}: {created.stderr.strip()}")
|
|
167
|
+
|
|
168
|
+
remote = github_manifest(manifest, repo, tag)
|
|
169
|
+
with tempfile.TemporaryDirectory() as tmp:
|
|
170
|
+
staging = Path(tmp)
|
|
171
|
+
uploads: list[str] = []
|
|
172
|
+
for ref in manifest.files():
|
|
173
|
+
if ref.path is None:
|
|
174
|
+
continue
|
|
175
|
+
target = staging / github_asset_name(manifest.id, ref.path)
|
|
176
|
+
if not target.exists():
|
|
177
|
+
shutil.copyfile(bundle.root / ref.path, target)
|
|
178
|
+
uploads.append(str(target))
|
|
179
|
+
manifest_file = staging / f"{manifest.id}.json"
|
|
180
|
+
manifest_file.write_text(remote.to_json(), encoding="utf-8")
|
|
181
|
+
uploads.append(str(manifest_file))
|
|
182
|
+
result = execute(["gh", "release", "upload", tag, "-R", repo, "--clobber", *uploads])
|
|
183
|
+
if result.returncode != 0:
|
|
184
|
+
raise ModelPortError(f"upload to {repo} {tag} failed: {result.stderr.strip()}")
|
|
185
|
+
|
|
186
|
+
url = f"https://github.com/{repo}/releases/download/{tag}/{manifest.id}.json"
|
|
187
|
+
return PublishResult(repo_id=repo, url=url, location=url, files=[Path(u).name for u in uploads])
|
|
188
|
+
|
|
189
|
+
|
|
190
|
+
def _run_gh(args: list[str]) -> subprocess.CompletedProcess[str]:
|
|
191
|
+
return subprocess.run(args, capture_output=True, text=True, check=False)
|
modelport/py.typed
ADDED
|
File without changes
|
modelport/quantize.py
ADDED
|
@@ -0,0 +1,120 @@
|
|
|
1
|
+
"""Smaller ONNX variants: fp16 weights, or int8 weights for MatMul and Gemm layers."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import math
|
|
6
|
+
import tempfile
|
|
7
|
+
from dataclasses import dataclass
|
|
8
|
+
from pathlib import Path
|
|
9
|
+
from typing import Literal
|
|
10
|
+
|
|
11
|
+
from .bundle import Bundle
|
|
12
|
+
from .errors import MissingDependencyError, ModelPortError
|
|
13
|
+
from .golden import golden_arrays
|
|
14
|
+
from .manifest import ClassificationPostprocess, Manifest, Runtime, Tolerance, Variant
|
|
15
|
+
from .runtimes import run_variant
|
|
16
|
+
from .verify import compare
|
|
17
|
+
|
|
18
|
+
Kind = Literal["fp16", "int8"]
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class QuantizeError(ModelPortError):
|
|
22
|
+
"""A smaller variant could not be made or is not accurate enough."""
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
@dataclass(frozen=True)
|
|
26
|
+
class QuantizeResult:
|
|
27
|
+
variant: Variant
|
|
28
|
+
size_ratio: float
|
|
29
|
+
max_abs_diff: float
|
|
30
|
+
top1_match: bool | None
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def tolerance_for(max_abs_diff: float, base: Tolerance) -> Tolerance:
|
|
34
|
+
"""Twice the measured error, rounded up to two significant digits, at least base.atol."""
|
|
35
|
+
limit = max(base.atol, 2 * max_abs_diff)
|
|
36
|
+
digits = 1 - math.floor(math.log10(limit))
|
|
37
|
+
return Tolerance(atol=math.ceil(limit * 10**digits) / 10**digits, rtol=base.rtol)
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def _write_variant(source: Path, target: Path, kind: Kind) -> None:
|
|
41
|
+
try:
|
|
42
|
+
import onnx
|
|
43
|
+
from onnxruntime.quantization import QuantType, quantize_dynamic
|
|
44
|
+
from onnxruntime.transformers.float16 import convert_float_to_float16
|
|
45
|
+
except ImportError as error:
|
|
46
|
+
raise MissingDependencyError("Quantizing ONNX models", "onnx") from error
|
|
47
|
+
|
|
48
|
+
model = onnx.load(str(source))
|
|
49
|
+
# Exporter shape annotations can disagree with ONNX shape inference; drop them.
|
|
50
|
+
del model.graph.value_info[:]
|
|
51
|
+
target.parent.mkdir(parents=True, exist_ok=True)
|
|
52
|
+
if kind == "fp16":
|
|
53
|
+
# Inputs and outputs stay float32 so apps feed the same tensors to every variant.
|
|
54
|
+
onnx.save(convert_float_to_float16(model, keep_io_types=True), str(target))
|
|
55
|
+
return
|
|
56
|
+
with tempfile.TemporaryDirectory() as tmp:
|
|
57
|
+
clean = Path(tmp) / "model.onnx"
|
|
58
|
+
onnx.save(model, str(clean))
|
|
59
|
+
# Dynamic int8 on convolutions badly hurts CNN accuracy, so only quantize
|
|
60
|
+
# MatMul and Gemm, where most transformer and classifier-head weights live.
|
|
61
|
+
quantize_dynamic(
|
|
62
|
+
str(clean),
|
|
63
|
+
str(target),
|
|
64
|
+
weight_type=QuantType.QInt8,
|
|
65
|
+
op_types_to_quantize=["MatMul", "Gemm"],
|
|
66
|
+
per_channel=True,
|
|
67
|
+
)
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
def quantize_bundle(
|
|
71
|
+
path: str | Path, kinds: list[Kind], *, allow_top1_change: bool = False
|
|
72
|
+
) -> list[QuantizeResult]:
|
|
73
|
+
"""Add one ONNX variant per kind, measure its error, and record a matching tolerance."""
|
|
74
|
+
bundle = Bundle(path)
|
|
75
|
+
manifest = bundle.read_manifest()
|
|
76
|
+
if manifest.golden is None:
|
|
77
|
+
raise QuantizeError("the bundle needs golden data to measure accuracy")
|
|
78
|
+
base = next(
|
|
79
|
+
(v for v in manifest.variants if v.runtime is Runtime.ONNX and v.precision == "fp32"),
|
|
80
|
+
None,
|
|
81
|
+
)
|
|
82
|
+
if base is None or base.file.path is None:
|
|
83
|
+
raise QuantizeError("the bundle has no onnx fp32 variant to start from")
|
|
84
|
+
|
|
85
|
+
inputs, expected = golden_arrays(bundle, manifest)
|
|
86
|
+
classify = {
|
|
87
|
+
o.name for o in manifest.outputs if isinstance(o.postprocess, ClassificationPostprocess)
|
|
88
|
+
}
|
|
89
|
+
variants = list(manifest.variants)
|
|
90
|
+
results = []
|
|
91
|
+
for kind in dict.fromkeys(kinds):
|
|
92
|
+
variant_id = f"onnx-{kind}"
|
|
93
|
+
relative = f"{variant_id}/model.onnx"
|
|
94
|
+
_write_variant(bundle.path(base.file.path), bundle.path(relative), kind)
|
|
95
|
+
candidate = Variant(
|
|
96
|
+
id=variant_id, runtime=Runtime.ONNX, precision=kind, file=bundle.add(relative)
|
|
97
|
+
)
|
|
98
|
+
got = run_variant(bundle, manifest, candidate, inputs)
|
|
99
|
+
checks = [
|
|
100
|
+
compare(name, got[name], want, manifest.golden.tolerance, name in classify)
|
|
101
|
+
for name, want in expected.items()
|
|
102
|
+
]
|
|
103
|
+
worst = max(check.max_abs_diff for check in checks)
|
|
104
|
+
top1 = None if not classify else all(c.top1_match is not False for c in checks)
|
|
105
|
+
if top1 is False and not allow_top1_change:
|
|
106
|
+
raise QuantizeError(
|
|
107
|
+
f"{variant_id} changes the top-1 class on the golden input "
|
|
108
|
+
f"(max |diff| {worst:.3g}). Use --allow-top1-change to keep it anyway."
|
|
109
|
+
)
|
|
110
|
+
variant = candidate.model_copy(
|
|
111
|
+
update={"tolerance": tolerance_for(worst, manifest.golden.tolerance)}
|
|
112
|
+
)
|
|
113
|
+
variants = [v for v in variants if v.id != variant_id] + [variant]
|
|
114
|
+
results.append(QuantizeResult(variant, variant.file.size / base.file.size, worst, top1))
|
|
115
|
+
|
|
116
|
+
updated = Manifest.model_validate(
|
|
117
|
+
{**manifest.model_dump(by_alias=True), "variants": [v.model_dump() for v in variants]}
|
|
118
|
+
)
|
|
119
|
+
bundle.write_manifest(updated)
|
|
120
|
+
return results
|
modelport/runtimes.py
ADDED
|
@@ -0,0 +1,52 @@
|
|
|
1
|
+
"""Run exported variants in Python, the same engines the Dart adapters wrap."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
|
|
7
|
+
from .bundle import Bundle
|
|
8
|
+
from .errors import MissingDependencyError, ModelPortError
|
|
9
|
+
from .manifest import Manifest, Runtime, Variant
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class RuntimeNotSupported(ModelPortError):
|
|
13
|
+
"""This runtime cannot run tensor inputs in Python."""
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def run_variant(
|
|
17
|
+
bundle: Bundle, manifest: Manifest, variant: Variant, inputs: dict[str, np.ndarray]
|
|
18
|
+
) -> dict[str, np.ndarray]:
|
|
19
|
+
"""Run one variant on named inputs and return outputs keyed by manifest output name."""
|
|
20
|
+
if variant.file.path is None:
|
|
21
|
+
raise ModelPortError(f"{variant.id}: only files inside the bundle can be run")
|
|
22
|
+
path = bundle.path(variant.file.path)
|
|
23
|
+
output_names = [o.name for o in manifest.outputs]
|
|
24
|
+
|
|
25
|
+
if variant.runtime is Runtime.ONNX:
|
|
26
|
+
try:
|
|
27
|
+
import onnxruntime as ort
|
|
28
|
+
except ImportError as error:
|
|
29
|
+
raise MissingDependencyError("Running ONNX variants", "onnx") from error
|
|
30
|
+
session = ort.InferenceSession(str(path), providers=["CPUExecutionProvider"])
|
|
31
|
+
results = session.run(output_names, inputs)
|
|
32
|
+
return {name: np.asarray(value) for name, value in zip(output_names, results, strict=True)}
|
|
33
|
+
|
|
34
|
+
if variant.runtime is Runtime.EXECUTORCH:
|
|
35
|
+
try:
|
|
36
|
+
import torch
|
|
37
|
+
from executorch.runtime import Runtime as ExecuTorchRuntime
|
|
38
|
+
except ImportError as error:
|
|
39
|
+
raise MissingDependencyError("Running ExecuTorch variants", "executorch") from error
|
|
40
|
+
method = ExecuTorchRuntime.get().load_program(str(path)).load_method("forward")
|
|
41
|
+
if method is None:
|
|
42
|
+
raise ModelPortError(f"{variant.id}: the program has no 'forward' method")
|
|
43
|
+
# ExecuTorch takes inputs by position, in manifest order.
|
|
44
|
+
args = [torch.from_numpy(np.ascontiguousarray(inputs[s.name])) for s in manifest.inputs]
|
|
45
|
+
results = method.execute(args)
|
|
46
|
+
return {
|
|
47
|
+
name: value.detach().numpy() for name, value in zip(output_names, results, strict=True)
|
|
48
|
+
}
|
|
49
|
+
|
|
50
|
+
raise RuntimeNotSupported(
|
|
51
|
+
f"{variant.id}: runtime '{variant.runtime}' has no tensor golden test"
|
|
52
|
+
)
|
|
@@ -0,0 +1,85 @@
|
|
|
1
|
+
"""Load a PyTorch model and describe it well enough to write a manifest.
|
|
2
|
+
|
|
3
|
+
Sources:
|
|
4
|
+
torchvision:<model name> e.g. torchvision:mobilenet_v3_small
|
|
5
|
+
hf:<repo id or local folder> e.g. hf:google/vit-base-patch16-224
|
|
6
|
+
file:<script.py>[:function] a function that returns a SourceModel
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
import re
|
|
12
|
+
from dataclasses import dataclass, field
|
|
13
|
+
from typing import Any
|
|
14
|
+
|
|
15
|
+
from ..errors import ModelPortError
|
|
16
|
+
from ..manifest import DetectionPostprocess, InputSpec, Task
|
|
17
|
+
|
|
18
|
+
_ID_CLEAN = re.compile(r"[^a-z0-9._-]+")
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class SourceError(ModelPortError):
|
|
22
|
+
"""A model source could not be loaded."""
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
@dataclass
|
|
26
|
+
class SourceModel:
|
|
27
|
+
"""A PyTorch module plus everything needed to describe it in a manifest.
|
|
28
|
+
|
|
29
|
+
`example_inputs` must match `inputs` in order and shape. They are made contiguous
|
|
30
|
+
before export, because ExecuTorch records the memory layout of example inputs.
|
|
31
|
+
"""
|
|
32
|
+
|
|
33
|
+
module: Any
|
|
34
|
+
example_inputs: tuple[Any, ...]
|
|
35
|
+
id: str
|
|
36
|
+
task: Task
|
|
37
|
+
license: str
|
|
38
|
+
inputs: list[InputSpec]
|
|
39
|
+
output_names: list[str]
|
|
40
|
+
labels: list[str] | None = None
|
|
41
|
+
name: str | None = None
|
|
42
|
+
description: str | None = None
|
|
43
|
+
source: str | None = None
|
|
44
|
+
top_k: int = 5
|
|
45
|
+
detection: DetectionPostprocess | None = None
|
|
46
|
+
"""For object detection: how to decode the outputs. Labels are attached on export."""
|
|
47
|
+
extra: dict[str, Any] = field(default_factory=dict)
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
def model_id(text: str) -> str:
|
|
51
|
+
"""Turn a model name or repo id into a manifest id, such as 'google/ViT-B' -> 'vit-b'."""
|
|
52
|
+
last = text.rstrip("/").split("/")[-1].lower()
|
|
53
|
+
cleaned = _ID_CLEAN.sub("-", last).strip("-._")
|
|
54
|
+
if not cleaned:
|
|
55
|
+
raise SourceError(f"cannot make a model id from '{text}'")
|
|
56
|
+
return cleaned
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def load_source(
|
|
60
|
+
spec: str, *, license: str | None = None, image_size: int | None = None
|
|
61
|
+
) -> SourceModel:
|
|
62
|
+
"""Load a model from a source string such as 'torchvision:mobilenet_v3_small'."""
|
|
63
|
+
kind, sep, rest = spec.partition(":")
|
|
64
|
+
if not sep or not rest:
|
|
65
|
+
raise SourceError(f"'{spec}' is not a source. Use torchvision:, hf:, or file:")
|
|
66
|
+
if kind == "torchvision":
|
|
67
|
+
from .torchvision import load_torchvision
|
|
68
|
+
|
|
69
|
+
model = load_torchvision(rest)
|
|
70
|
+
elif kind == "hf":
|
|
71
|
+
from .hf import load_hf
|
|
72
|
+
|
|
73
|
+
model = load_hf(rest, license=license, image_size=image_size)
|
|
74
|
+
elif kind == "file":
|
|
75
|
+
from .file import load_file
|
|
76
|
+
|
|
77
|
+
model = load_file(rest)
|
|
78
|
+
else:
|
|
79
|
+
raise SourceError(f"unknown source type '{kind}'. Use torchvision:, hf:, or file:")
|
|
80
|
+
if license is not None:
|
|
81
|
+
model.license = license
|
|
82
|
+
return model
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
__all__ = ["SourceError", "SourceModel", "load_source", "model_id"]
|
|
@@ -0,0 +1,46 @@
|
|
|
1
|
+
"""Models described by a function in the user's own Python file."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import importlib.util
|
|
6
|
+
from pathlib import Path
|
|
7
|
+
|
|
8
|
+
from . import SourceError, SourceModel
|
|
9
|
+
|
|
10
|
+
DEFAULT_FUNCTION = "build"
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def split_file_spec(spec: str) -> tuple[Path, str]:
|
|
14
|
+
"""'models/net.py:make' -> (models/net.py, 'make'). The function defaults to 'build'."""
|
|
15
|
+
path_text, sep, function = spec.rpartition(":")
|
|
16
|
+
if sep and path_text.endswith(".py") and function:
|
|
17
|
+
return Path(path_text), function
|
|
18
|
+
return Path(spec), DEFAULT_FUNCTION
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def load_file(spec: str) -> SourceModel:
|
|
22
|
+
"""Import a Python file and call a function that returns a SourceModel.
|
|
23
|
+
|
|
24
|
+
This runs the file's code, like `python file.py` would. Only use files you trust.
|
|
25
|
+
"""
|
|
26
|
+
path, function = split_file_spec(spec)
|
|
27
|
+
if not path.is_file():
|
|
28
|
+
raise SourceError(f"{path} does not exist")
|
|
29
|
+
module_spec = importlib.util.spec_from_file_location(f"_modelport_user_{path.stem}", path)
|
|
30
|
+
if module_spec is None or module_spec.loader is None:
|
|
31
|
+
raise SourceError(f"cannot import {path}")
|
|
32
|
+
module = importlib.util.module_from_spec(module_spec)
|
|
33
|
+
module_spec.loader.exec_module(module)
|
|
34
|
+
|
|
35
|
+
factory = getattr(module, function, None)
|
|
36
|
+
if not callable(factory):
|
|
37
|
+
raise SourceError(f"{path} has no function named '{function}'")
|
|
38
|
+
result = factory()
|
|
39
|
+
if not isinstance(result, SourceModel):
|
|
40
|
+
raise SourceError(
|
|
41
|
+
f"{path}:{function} must return modelport.sources.SourceModel, "
|
|
42
|
+
f"got {type(result).__name__}"
|
|
43
|
+
)
|
|
44
|
+
if result.source is None:
|
|
45
|
+
result.source = f"file:{path.name}:{function}"
|
|
46
|
+
return result
|