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.
@@ -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