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
|
@@ -0,0 +1,227 @@
|
|
|
1
|
+
"""Generate a typed Dart wrapper for a tensor model's inputs and outputs.
|
|
2
|
+
|
|
3
|
+
The wrapper turns tensor names into named parameters and fields, so a typo in an
|
|
4
|
+
input name is a compile error instead of a runtime one.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
import re
|
|
10
|
+
|
|
11
|
+
from ..errors import ModelPortError
|
|
12
|
+
from ..manifest import InputSpec, Manifest, OutputSpec, Task, TensorSpec
|
|
13
|
+
|
|
14
|
+
# Dart keywords and built-in identifiers that cannot be parameter or field names.
|
|
15
|
+
_RESERVED = set(
|
|
16
|
+
[
|
|
17
|
+
"abstract",
|
|
18
|
+
"as",
|
|
19
|
+
"assert",
|
|
20
|
+
"async",
|
|
21
|
+
"await",
|
|
22
|
+
"base",
|
|
23
|
+
"break",
|
|
24
|
+
"case",
|
|
25
|
+
"catch",
|
|
26
|
+
"class",
|
|
27
|
+
"const",
|
|
28
|
+
"continue",
|
|
29
|
+
"covariant",
|
|
30
|
+
"default",
|
|
31
|
+
"deferred",
|
|
32
|
+
"do",
|
|
33
|
+
"dynamic",
|
|
34
|
+
"else",
|
|
35
|
+
"enum",
|
|
36
|
+
"export",
|
|
37
|
+
"extends",
|
|
38
|
+
"extension",
|
|
39
|
+
"external",
|
|
40
|
+
"factory",
|
|
41
|
+
"false",
|
|
42
|
+
"final",
|
|
43
|
+
"finally",
|
|
44
|
+
"for",
|
|
45
|
+
"function",
|
|
46
|
+
"get",
|
|
47
|
+
"hide",
|
|
48
|
+
"if",
|
|
49
|
+
"implements",
|
|
50
|
+
"import",
|
|
51
|
+
"in",
|
|
52
|
+
"interface",
|
|
53
|
+
"is",
|
|
54
|
+
"late",
|
|
55
|
+
"library",
|
|
56
|
+
"mixin",
|
|
57
|
+
"new",
|
|
58
|
+
"null",
|
|
59
|
+
"of",
|
|
60
|
+
"on",
|
|
61
|
+
"operator",
|
|
62
|
+
"part",
|
|
63
|
+
"required",
|
|
64
|
+
"rethrow",
|
|
65
|
+
"return",
|
|
66
|
+
"sealed",
|
|
67
|
+
"set",
|
|
68
|
+
"show",
|
|
69
|
+
"static",
|
|
70
|
+
"super",
|
|
71
|
+
"switch",
|
|
72
|
+
"sync",
|
|
73
|
+
"this",
|
|
74
|
+
"throw",
|
|
75
|
+
"true",
|
|
76
|
+
"try",
|
|
77
|
+
"type",
|
|
78
|
+
"typedef",
|
|
79
|
+
"var",
|
|
80
|
+
"void",
|
|
81
|
+
"when",
|
|
82
|
+
"while",
|
|
83
|
+
"with",
|
|
84
|
+
"yield",
|
|
85
|
+
]
|
|
86
|
+
)
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
def _words(text: str) -> list[str]:
|
|
90
|
+
return [w for w in re.split(r"[^A-Za-z0-9]+", text) if w]
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
def class_name(model_id: str) -> str:
|
|
94
|
+
"""'mobilenet_v3_small' -> 'MobilenetV3Small'."""
|
|
95
|
+
name = "".join(w[:1].upper() + w[1:] for w in _words(model_id))
|
|
96
|
+
return name if name[:1].isalpha() else f"Model{name}"
|
|
97
|
+
|
|
98
|
+
|
|
99
|
+
def member_name(tensor: str) -> str:
|
|
100
|
+
"""'pixel_values' -> 'pixelValues'."""
|
|
101
|
+
words = _words(tensor) or ["value"]
|
|
102
|
+
name = words[0][:1].lower() + words[0][1:] + "".join(w[:1].upper() + w[1:] for w in words[1:])
|
|
103
|
+
if not name[:1].isalpha():
|
|
104
|
+
name = f"t{name}"
|
|
105
|
+
return f"{name}_" if name in _RESERVED else name
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
def file_name(model_id: str) -> str:
|
|
109
|
+
"""'qwen2.5-0.5b' -> 'qwen2_5_0_5b.dart'."""
|
|
110
|
+
return "_".join(w.lower() for w in _words(model_id)) + ".dart"
|
|
111
|
+
|
|
112
|
+
|
|
113
|
+
def _doc(spec: TensorSpec) -> str:
|
|
114
|
+
# Backticks keep dartdoc from reading the shape as a reference.
|
|
115
|
+
return f"`{spec.name}`: {spec.dtype.value} `{spec.shape}`"
|
|
116
|
+
|
|
117
|
+
|
|
118
|
+
def _unique(names: list[str], what: str) -> None:
|
|
119
|
+
if len(set(names)) != len(names):
|
|
120
|
+
raise ModelPortError(f"two {what} names map to the same Dart name: {names}")
|
|
121
|
+
|
|
122
|
+
|
|
123
|
+
def generate_dart(manifest: Manifest, *, location: str | None = None) -> str:
|
|
124
|
+
"""Dart source for a typed wrapper around `TensorModel`."""
|
|
125
|
+
if manifest.task is Task.TEXT_GENERATION:
|
|
126
|
+
raise ModelPortError("text-generation models have no tensors; use TextGenerator in Dart")
|
|
127
|
+
name = class_name(manifest.id)
|
|
128
|
+
inputs: list[InputSpec] = list(manifest.inputs)
|
|
129
|
+
outputs: list[OutputSpec] = list(manifest.outputs)
|
|
130
|
+
in_names = [member_name(i.name) for i in inputs]
|
|
131
|
+
out_names = [member_name(o.name) for o in outputs]
|
|
132
|
+
_unique(in_names, "input")
|
|
133
|
+
_unique(out_names, "output")
|
|
134
|
+
|
|
135
|
+
lines: list[str] = [
|
|
136
|
+
f"// Generated by `modelport gen-dart` from {manifest.id} {manifest.version}.",
|
|
137
|
+
"// Do not edit by hand; run gen-dart again after the manifest changes.",
|
|
138
|
+
"",
|
|
139
|
+
"import 'package:modelport/modelport.dart';",
|
|
140
|
+
"",
|
|
141
|
+
f"/// Typed access to {manifest.name or manifest.id} ({manifest.task.value}).",
|
|
142
|
+
]
|
|
143
|
+
if manifest.description:
|
|
144
|
+
lines.append(f"///\n/// {manifest.description}")
|
|
145
|
+
lines += [f"class {name}Model {{", f" {name}Model._(this.model);", ""]
|
|
146
|
+
if location is not None:
|
|
147
|
+
lines += [
|
|
148
|
+
" /// Where the bundle is published.",
|
|
149
|
+
f" static const defaultLocation = '{location}';",
|
|
150
|
+
"",
|
|
151
|
+
]
|
|
152
|
+
for spec, member in zip(inputs, in_names, strict=True):
|
|
153
|
+
lines += [
|
|
154
|
+
f" /// Shape of input {_doc(spec)}.",
|
|
155
|
+
f" static const {member}Shape = <int>{spec.shape};",
|
|
156
|
+
"",
|
|
157
|
+
]
|
|
158
|
+
|
|
159
|
+
lines.append(" /// Downloads (if needed), verifies, and opens the model.")
|
|
160
|
+
if location is not None:
|
|
161
|
+
lines += [f" static Future<{name}Model> load({{", " String location = defaultLocation,"]
|
|
162
|
+
else:
|
|
163
|
+
lines += [f" static Future<{name}Model> load(", " String location, {"]
|
|
164
|
+
lines += [
|
|
165
|
+
" String? variantId,",
|
|
166
|
+
" void Function(DownloadProgress)? onProgress,",
|
|
167
|
+
" }) async =>",
|
|
168
|
+
f" {name}Model._(",
|
|
169
|
+
" await ModelPort.load(location, variantId: variantId, onProgress: onProgress),",
|
|
170
|
+
" );",
|
|
171
|
+
"",
|
|
172
|
+
" /// The underlying model, for golden checks and raw access.",
|
|
173
|
+
" final TensorModel model;",
|
|
174
|
+
"",
|
|
175
|
+
" /// Runs the model.",
|
|
176
|
+
" ///",
|
|
177
|
+
]
|
|
178
|
+
for spec in inputs:
|
|
179
|
+
lines.append(f" /// * {_doc(spec)}")
|
|
180
|
+
params = "".join(f"required Tensor {m}, " for m in in_names).rstrip(", ")
|
|
181
|
+
lines += [
|
|
182
|
+
f" Future<{name}Outputs> run({{{params}}}) async {{",
|
|
183
|
+
" final outputs = await model.run({",
|
|
184
|
+
]
|
|
185
|
+
for spec, member in zip(inputs, in_names, strict=True):
|
|
186
|
+
lines.append(f" '{spec.name}': {member},")
|
|
187
|
+
lines += [" });", f" return {name}Outputs("]
|
|
188
|
+
for spec, member in zip(outputs, out_names, strict=True):
|
|
189
|
+
lines.append(f" {member}: outputs['{spec.name}']!,")
|
|
190
|
+
lines += [" );", " }", "", " Future<void> close() => model.close();", "}", ""]
|
|
191
|
+
|
|
192
|
+
lines += [
|
|
193
|
+
f"/// Outputs of {name}Model.run.",
|
|
194
|
+
f"class {name}Outputs {{",
|
|
195
|
+
f" const {name}Outputs({{",
|
|
196
|
+
]
|
|
197
|
+
for member in out_names:
|
|
198
|
+
lines.append(f" required this.{member},")
|
|
199
|
+
lines += [" });", ""]
|
|
200
|
+
for spec, member in zip(outputs, out_names, strict=True):
|
|
201
|
+
lines += [f" /// {_doc(spec)}.", f" final Tensor {member};", ""]
|
|
202
|
+
lines[-1] = "}"
|
|
203
|
+
return "\n".join(lines) + "\n"
|
|
204
|
+
|
|
205
|
+
|
|
206
|
+
def format_with_dart(source: str) -> str | None:
|
|
207
|
+
"""Format Dart source with `dart format`, or return None if Dart is not installed."""
|
|
208
|
+
import shutil
|
|
209
|
+
import subprocess
|
|
210
|
+
import tempfile
|
|
211
|
+
from pathlib import Path
|
|
212
|
+
|
|
213
|
+
dart = shutil.which("dart")
|
|
214
|
+
if dart is None:
|
|
215
|
+
return None
|
|
216
|
+
with tempfile.TemporaryDirectory() as tmp:
|
|
217
|
+
path = Path(tmp) / "generated.dart"
|
|
218
|
+
path.write_text(source, encoding="utf-8")
|
|
219
|
+
result = subprocess.run(
|
|
220
|
+
[dart, "format", "--output=write", str(path)],
|
|
221
|
+
capture_output=True,
|
|
222
|
+
text=True,
|
|
223
|
+
check=False,
|
|
224
|
+
)
|
|
225
|
+
if result.returncode != 0:
|
|
226
|
+
raise ModelPortError(f"dart format failed on generated code:\n{result.stderr}")
|
|
227
|
+
return path.read_text(encoding="utf-8")
|
modelport/doctor.py
ADDED
|
@@ -0,0 +1,88 @@
|
|
|
1
|
+
"""Check the local environment for what each CLI feature needs."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import platform
|
|
6
|
+
import shutil
|
|
7
|
+
import sys
|
|
8
|
+
from collections.abc import Callable
|
|
9
|
+
from dataclasses import dataclass, field
|
|
10
|
+
from importlib import metadata
|
|
11
|
+
from pathlib import Path
|
|
12
|
+
|
|
13
|
+
from . import __version__
|
|
14
|
+
|
|
15
|
+
# extra name -> packages it installs, in the order shown to users
|
|
16
|
+
EXTRAS: dict[str, tuple[str, ...]] = {
|
|
17
|
+
"onnx": ("onnx", "onnxruntime", "onnxscript", "torch"),
|
|
18
|
+
"executorch": ("executorch", "torch"),
|
|
19
|
+
"gguf": ("gguf",),
|
|
20
|
+
"torchvision": ("torchvision", "torch"),
|
|
21
|
+
"hf": ("transformers", "huggingface-hub", "torch"),
|
|
22
|
+
}
|
|
23
|
+
|
|
24
|
+
LOW_DISK_BYTES = 5 * 1000**3
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
@dataclass(frozen=True)
|
|
28
|
+
class PackageStatus:
|
|
29
|
+
name: str
|
|
30
|
+
extra: str
|
|
31
|
+
version: str | None
|
|
32
|
+
|
|
33
|
+
@property
|
|
34
|
+
def installed(self) -> bool:
|
|
35
|
+
return self.version is not None
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
@dataclass(frozen=True)
|
|
39
|
+
class DoctorReport:
|
|
40
|
+
modelport_version: str
|
|
41
|
+
python_version: str
|
|
42
|
+
platform: str
|
|
43
|
+
free_disk_bytes: int
|
|
44
|
+
packages: list[PackageStatus]
|
|
45
|
+
problems: list[str] = field(default_factory=list)
|
|
46
|
+
hints: list[str] = field(default_factory=list)
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def installed_version(name: str) -> str | None:
|
|
50
|
+
try:
|
|
51
|
+
return metadata.version(name)
|
|
52
|
+
except metadata.PackageNotFoundError:
|
|
53
|
+
return None
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
def collect_report(
|
|
57
|
+
folder: Path | None = None,
|
|
58
|
+
version_of: Callable[[str], str | None] = installed_version,
|
|
59
|
+
) -> DoctorReport:
|
|
60
|
+
"""Gather versions and disk space. Never imports heavy packages such as torch."""
|
|
61
|
+
packages = [
|
|
62
|
+
PackageStatus(name=name, extra=extra, version=version_of(name))
|
|
63
|
+
for extra, names in EXTRAS.items()
|
|
64
|
+
for name in names
|
|
65
|
+
]
|
|
66
|
+
free = shutil.disk_usage(folder or Path.cwd()).free
|
|
67
|
+
|
|
68
|
+
problems: list[str] = []
|
|
69
|
+
hints: list[str] = []
|
|
70
|
+
for extra, names in EXTRAS.items():
|
|
71
|
+
statuses = [p for p in packages if p.name in names]
|
|
72
|
+
missing = [p.name for p in statuses if not p.installed]
|
|
73
|
+
if missing and len(missing) < len(statuses):
|
|
74
|
+
problems.append(f"'{extra}' is only partly installed, missing: {', '.join(missing)}")
|
|
75
|
+
if missing:
|
|
76
|
+
hints.append(f"pip install 'modelport-cli[{extra}]'")
|
|
77
|
+
if free < LOW_DISK_BYTES:
|
|
78
|
+
problems.append(f"Only {free / 1000**3:.1f} GB free. Exports and LLM downloads need more.")
|
|
79
|
+
|
|
80
|
+
return DoctorReport(
|
|
81
|
+
modelport_version=__version__,
|
|
82
|
+
python_version=platform.python_version(),
|
|
83
|
+
platform=f"{platform.system()} {platform.machine()} ({sys.platform})",
|
|
84
|
+
free_disk_bytes=free,
|
|
85
|
+
packages=packages,
|
|
86
|
+
problems=problems,
|
|
87
|
+
hints=hints,
|
|
88
|
+
)
|
modelport/errors.py
ADDED
|
@@ -0,0 +1,18 @@
|
|
|
1
|
+
"""Errors that the CLI turns into short messages instead of tracebacks."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class ModelPortError(Exception):
|
|
7
|
+
"""Base class for expected, user-facing errors."""
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class MissingDependencyError(ModelPortError):
|
|
11
|
+
"""An optional package needed for this feature is not installed."""
|
|
12
|
+
|
|
13
|
+
def __init__(self, feature: str, extra: str) -> None:
|
|
14
|
+
super().__init__(
|
|
15
|
+
f"{feature} needs extra packages. "
|
|
16
|
+
f"Install them with: pip install 'modelport-cli[{extra}]'"
|
|
17
|
+
)
|
|
18
|
+
self.extra = extra
|
|
@@ -0,0 +1,55 @@
|
|
|
1
|
+
"""Turn a SourceModel into model files for each target runtime."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import contextlib
|
|
6
|
+
import io
|
|
7
|
+
from collections.abc import Callable, Iterator
|
|
8
|
+
from dataclasses import dataclass, field
|
|
9
|
+
|
|
10
|
+
from ..bundle import Bundle
|
|
11
|
+
from ..errors import ModelPortError
|
|
12
|
+
from ..manifest import Runtime
|
|
13
|
+
from ..sources import SourceModel
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
class ExportError(ModelPortError):
|
|
17
|
+
"""A model could not be exported."""
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
@dataclass(frozen=True)
|
|
21
|
+
class ExportedVariant:
|
|
22
|
+
"""Files written for one variant, as bundle-relative paths."""
|
|
23
|
+
|
|
24
|
+
id: str
|
|
25
|
+
runtime: Runtime
|
|
26
|
+
precision: str
|
|
27
|
+
main_file: str
|
|
28
|
+
extra_files: list[str] = field(default_factory=list)
|
|
29
|
+
backend: str | None = None
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
Exporter = Callable[[SourceModel, Bundle], ExportedVariant]
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
@contextlib.contextmanager
|
|
36
|
+
def quiet_output() -> Iterator[io.StringIO]:
|
|
37
|
+
"""Capture the progress lines that exporters print, to show them only on failure."""
|
|
38
|
+
buffer = io.StringIO()
|
|
39
|
+
with contextlib.redirect_stdout(buffer), contextlib.redirect_stderr(buffer):
|
|
40
|
+
yield buffer
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
def get_exporter(target: str) -> Exporter:
|
|
44
|
+
if target == "onnx":
|
|
45
|
+
from .onnx import export_onnx
|
|
46
|
+
|
|
47
|
+
return export_onnx
|
|
48
|
+
if target == "executorch":
|
|
49
|
+
from .executorch import export_executorch
|
|
50
|
+
|
|
51
|
+
return export_executorch
|
|
52
|
+
raise ExportError(f"unknown export target '{target}'. Available: onnx, executorch")
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
__all__ = ["ExportError", "ExportedVariant", "Exporter", "get_exporter", "quiet_output"]
|
|
@@ -0,0 +1,46 @@
|
|
|
1
|
+
"""Export to an ExecuTorch program (.pte) delegated to XNNPACK."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from ..bundle import Bundle
|
|
6
|
+
from ..errors import MissingDependencyError
|
|
7
|
+
from ..manifest import Runtime
|
|
8
|
+
from ..sources import SourceModel
|
|
9
|
+
from . import ExportedVariant, ExportError, quiet_output
|
|
10
|
+
|
|
11
|
+
VARIANT_ID = "executorch-xnnpack-fp32"
|
|
12
|
+
MODEL_PATH = f"{VARIANT_ID}/model.pte"
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def export_executorch(source: SourceModel, bundle: Bundle) -> ExportedVariant:
|
|
16
|
+
try:
|
|
17
|
+
import torch
|
|
18
|
+
from executorch.backends.xnnpack.partition.xnnpack_partitioner import (
|
|
19
|
+
XnnpackPartitioner,
|
|
20
|
+
)
|
|
21
|
+
from executorch.exir import to_edge_transform_and_lower
|
|
22
|
+
except ImportError as error:
|
|
23
|
+
raise MissingDependencyError("Exporting to ExecuTorch", "executorch") from error
|
|
24
|
+
|
|
25
|
+
# ExecuTorch records the memory layout of example inputs. A channels_last example
|
|
26
|
+
# makes the program reject ordinary NCHW tensors at run time (found in Phase 0).
|
|
27
|
+
example = tuple(t.contiguous() for t in source.example_inputs)
|
|
28
|
+
with quiet_output() as log, torch.no_grad():
|
|
29
|
+
try:
|
|
30
|
+
exported = torch.export.export(source.module, example)
|
|
31
|
+
program = to_edge_transform_and_lower(
|
|
32
|
+
exported, partitioner=[XnnpackPartitioner()]
|
|
33
|
+
).to_executorch()
|
|
34
|
+
except Exception as error:
|
|
35
|
+
raise ExportError(f"ExecuTorch export failed: {error}\n{log.getvalue()}") from error
|
|
36
|
+
|
|
37
|
+
target = bundle.path(MODEL_PATH)
|
|
38
|
+
target.parent.mkdir(parents=True, exist_ok=True)
|
|
39
|
+
target.write_bytes(program.buffer)
|
|
40
|
+
return ExportedVariant(
|
|
41
|
+
id=VARIANT_ID,
|
|
42
|
+
runtime=Runtime.EXECUTORCH,
|
|
43
|
+
precision="fp32",
|
|
44
|
+
backend="xnnpack",
|
|
45
|
+
main_file=MODEL_PATH,
|
|
46
|
+
)
|
|
@@ -0,0 +1,63 @@
|
|
|
1
|
+
"""Export to ONNX with the torch.export-based (dynamo) exporter."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from ..bundle import Bundle
|
|
6
|
+
from ..errors import MissingDependencyError
|
|
7
|
+
from ..manifest import Runtime
|
|
8
|
+
from ..sources import SourceModel
|
|
9
|
+
from . import ExportedVariant, ExportError, quiet_output
|
|
10
|
+
|
|
11
|
+
VARIANT_ID = "onnx-fp32"
|
|
12
|
+
MODEL_PATH = f"{VARIANT_ID}/model.onnx"
|
|
13
|
+
# ONNX protobuf files cannot exceed 2 GB, so larger weights go to a side file.
|
|
14
|
+
EXTERNAL_DATA_BYTES = 1_800_000_000
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def export_onnx(source: SourceModel, bundle: Bundle) -> ExportedVariant:
|
|
18
|
+
try:
|
|
19
|
+
import onnx
|
|
20
|
+
import torch
|
|
21
|
+
except ImportError as error:
|
|
22
|
+
raise MissingDependencyError("Exporting to ONNX", "onnx") from error
|
|
23
|
+
|
|
24
|
+
target = bundle.path(MODEL_PATH)
|
|
25
|
+
target.parent.mkdir(parents=True, exist_ok=True)
|
|
26
|
+
weight_bytes = sum(p.numel() * p.element_size() for p in source.module.parameters())
|
|
27
|
+
external = weight_bytes > EXTERNAL_DATA_BYTES
|
|
28
|
+
# Contiguous inputs keep the exported graph in plain NCHW memory layout.
|
|
29
|
+
example = tuple(t.contiguous() for t in source.example_inputs)
|
|
30
|
+
|
|
31
|
+
with quiet_output() as log, torch.no_grad():
|
|
32
|
+
try:
|
|
33
|
+
torch.onnx.export(
|
|
34
|
+
source.module,
|
|
35
|
+
example,
|
|
36
|
+
str(target),
|
|
37
|
+
dynamo=True,
|
|
38
|
+
input_names=[spec.name for spec in source.inputs],
|
|
39
|
+
output_names=source.output_names,
|
|
40
|
+
external_data=external,
|
|
41
|
+
)
|
|
42
|
+
except Exception as error:
|
|
43
|
+
raise ExportError(f"ONNX export failed: {error}\n{log.getvalue()}") from error
|
|
44
|
+
|
|
45
|
+
onnx.checker.check_model(str(target))
|
|
46
|
+
graph = onnx.load(str(target), load_external_data=False).graph
|
|
47
|
+
initializers = {init.name for init in graph.initializer}
|
|
48
|
+
got_inputs = [v.name for v in graph.input if v.name not in initializers]
|
|
49
|
+
got_outputs = [v.name for v in graph.output]
|
|
50
|
+
if got_inputs != [s.name for s in source.inputs] or got_outputs != source.output_names:
|
|
51
|
+
raise ExportError(
|
|
52
|
+
f"exported names {got_inputs} -> {got_outputs} do not match "
|
|
53
|
+
f"{[s.name for s in source.inputs]} -> {source.output_names}"
|
|
54
|
+
)
|
|
55
|
+
|
|
56
|
+
extra = [f"{MODEL_PATH}.data"] if bundle.path(f"{MODEL_PATH}.data").is_file() else []
|
|
57
|
+
return ExportedVariant(
|
|
58
|
+
id=VARIANT_ID,
|
|
59
|
+
runtime=Runtime.ONNX,
|
|
60
|
+
precision="fp32",
|
|
61
|
+
main_file=MODEL_PATH,
|
|
62
|
+
extra_files=extra,
|
|
63
|
+
)
|
modelport/gguf_import.py
ADDED
|
@@ -0,0 +1,197 @@
|
|
|
1
|
+
"""Describe ready-made GGUF language models with a manifest.
|
|
2
|
+
|
|
3
|
+
Large GGUF files are not copied: the manifest points at them with pinned
|
|
4
|
+
Hugging Face URLs, so a bundle is just modelport.json.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
import math
|
|
10
|
+
import re
|
|
11
|
+
import shutil
|
|
12
|
+
from dataclasses import dataclass, field
|
|
13
|
+
from pathlib import Path
|
|
14
|
+
from typing import Any
|
|
15
|
+
|
|
16
|
+
from .bundle import Bundle
|
|
17
|
+
from .errors import MissingDependencyError, ModelPortError
|
|
18
|
+
from .manifest import FileRef, LlmConfig, Manifest, Runtime, Task, Variant
|
|
19
|
+
from .sources import model_id as make_id
|
|
20
|
+
|
|
21
|
+
_QUANT = re.compile(r"[-_.]((?:i?q\d\w*)|bf16|f16|f32)\.gguf$", re.IGNORECASE)
|
|
22
|
+
_SPLIT = re.compile(r"-\d{5}-of-\d{5}\.gguf$", re.IGNORECASE)
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
class GgufImportError(ModelPortError):
|
|
26
|
+
"""A GGUF model could not be described."""
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
@dataclass(frozen=True)
|
|
30
|
+
class GgufImport:
|
|
31
|
+
manifest: Manifest
|
|
32
|
+
warnings: list[str] = field(default_factory=list)
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def quant_of(filename: str) -> str | None:
|
|
36
|
+
"""'qwen2.5-0.5b-instruct-q4_k_m.gguf' -> 'q4_k_m'."""
|
|
37
|
+
match = _QUANT.search(filename)
|
|
38
|
+
return match.group(1).lower() if match else None
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def estimate_min_ram_mb(size_bytes: int, context: int) -> int:
|
|
42
|
+
"""A rough lower bound: weights, a KV cache for `context`, and runtime overhead."""
|
|
43
|
+
mib = size_bytes / 2**20 * 1.1 + 256 + context / 1024 * 32
|
|
44
|
+
return int(math.ceil(mib / 256) * 256)
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def _bundle_id(name: str) -> str:
|
|
48
|
+
base = make_id(name)
|
|
49
|
+
return base[: -len("-gguf")] if base.endswith("-gguf") else base
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def import_from_hub(
|
|
53
|
+
repo: str,
|
|
54
|
+
quants: list[str],
|
|
55
|
+
*,
|
|
56
|
+
revision: str | None = None,
|
|
57
|
+
context: int = 4096,
|
|
58
|
+
license: str | None = None,
|
|
59
|
+
model_id: str | None = None,
|
|
60
|
+
) -> GgufImport:
|
|
61
|
+
"""Build a manifest for GGUF files in a Hugging Face repo, without downloading them."""
|
|
62
|
+
try:
|
|
63
|
+
from huggingface_hub import HfApi
|
|
64
|
+
except ImportError as error:
|
|
65
|
+
raise MissingDependencyError("Importing GGUF models from the Hub", "hf") from error
|
|
66
|
+
|
|
67
|
+
api = HfApi()
|
|
68
|
+
try:
|
|
69
|
+
info = api.model_info(repo, revision=revision, files_metadata=True)
|
|
70
|
+
summary = getattr(api.model_info(repo, revision=revision, expand=["gguf"]), "gguf", None)
|
|
71
|
+
except Exception as error:
|
|
72
|
+
raise GgufImportError(
|
|
73
|
+
f"could not read {repo} from the Hugging Face Hub: {error}"
|
|
74
|
+
) from error
|
|
75
|
+
|
|
76
|
+
files: dict[str, Any] = {}
|
|
77
|
+
for sibling in info.siblings or []:
|
|
78
|
+
name = sibling.rfilename
|
|
79
|
+
if not name.endswith(".gguf") or _SPLIT.search(name):
|
|
80
|
+
continue
|
|
81
|
+
quant = quant_of(Path(name).name)
|
|
82
|
+
if quant and quant not in files:
|
|
83
|
+
files[quant] = sibling
|
|
84
|
+
if not files:
|
|
85
|
+
raise GgufImportError(f"{repo} has no single-file GGUF models")
|
|
86
|
+
|
|
87
|
+
warnings: list[str] = []
|
|
88
|
+
gguf = summary if isinstance(summary, dict) else {}
|
|
89
|
+
trained = gguf.get("context_length")
|
|
90
|
+
if trained and context > trained:
|
|
91
|
+
warnings.append(f"context {context} is above the trained {trained}; using {trained}")
|
|
92
|
+
context = int(trained)
|
|
93
|
+
if not gguf.get("chat_template"):
|
|
94
|
+
warnings.append("the GGUF has no chat template, so chat may not work")
|
|
95
|
+
|
|
96
|
+
license = license or _card_license(info)
|
|
97
|
+
if not license:
|
|
98
|
+
raise GgufImportError(f"no license found for {repo}. Pass it with --license")
|
|
99
|
+
|
|
100
|
+
variants = []
|
|
101
|
+
for quant in [q.lower() for q in quants]:
|
|
102
|
+
sibling = files.get(quant)
|
|
103
|
+
if sibling is None:
|
|
104
|
+
raise GgufImportError(
|
|
105
|
+
f"{repo} has no {quant} file. Available: {', '.join(sorted(files))}"
|
|
106
|
+
)
|
|
107
|
+
sha256 = _lfs_sha256(sibling)
|
|
108
|
+
if sha256 is None or not sibling.size:
|
|
109
|
+
raise GgufImportError(f"{sibling.rfilename} has no size or sha256 on the Hub")
|
|
110
|
+
url = f"https://huggingface.co/{repo}/resolve/{info.sha}/{sibling.rfilename}"
|
|
111
|
+
variants.append(
|
|
112
|
+
Variant(
|
|
113
|
+
id=f"gguf-{quant}",
|
|
114
|
+
runtime=Runtime.LLAMACPP,
|
|
115
|
+
precision=quant,
|
|
116
|
+
file=FileRef(url=url, size=sibling.size, sha256=sha256),
|
|
117
|
+
min_ram_mb=estimate_min_ram_mb(sibling.size, context),
|
|
118
|
+
)
|
|
119
|
+
)
|
|
120
|
+
|
|
121
|
+
manifest = Manifest(
|
|
122
|
+
id=model_id or _bundle_id(repo),
|
|
123
|
+
version="1.0.0",
|
|
124
|
+
task=Task.TEXT_GENERATION,
|
|
125
|
+
license=license,
|
|
126
|
+
name=repo.split("/")[-1],
|
|
127
|
+
description=f"{gguf.get('architecture', 'GGUF')} model from {repo}.",
|
|
128
|
+
source=f"hf:{repo}@{(info.sha or '')[:12]}",
|
|
129
|
+
variants=variants,
|
|
130
|
+
llm=LlmConfig(context_length=context),
|
|
131
|
+
)
|
|
132
|
+
return GgufImport(manifest, warnings)
|
|
133
|
+
|
|
134
|
+
|
|
135
|
+
def import_from_file(
|
|
136
|
+
path: str | Path,
|
|
137
|
+
bundle: Bundle,
|
|
138
|
+
*,
|
|
139
|
+
context: int = 4096,
|
|
140
|
+
license: str,
|
|
141
|
+
model_id: str | None = None,
|
|
142
|
+
) -> GgufImport:
|
|
143
|
+
"""Copy a local GGUF file into a bundle and describe it."""
|
|
144
|
+
from .inspection import inspect_model
|
|
145
|
+
|
|
146
|
+
source = Path(path)
|
|
147
|
+
info = inspect_model(source)
|
|
148
|
+
if info.format != "gguf":
|
|
149
|
+
raise GgufImportError(f"{source} is not a GGUF file")
|
|
150
|
+
quant = quant_of(source.name) or info.metadata.get("quantization", "unknown").lower()
|
|
151
|
+
warnings: list[str] = []
|
|
152
|
+
trained = info.metadata.get("context_length", "")
|
|
153
|
+
if trained.isdigit() and context > int(trained):
|
|
154
|
+
warnings.append(f"context {context} is above the trained {trained}; using {trained}")
|
|
155
|
+
context = int(trained)
|
|
156
|
+
if info.metadata.get("chat_template") != "yes":
|
|
157
|
+
warnings.append("the GGUF has no chat template, so chat may not work")
|
|
158
|
+
|
|
159
|
+
variant_id = f"gguf-{quant}"
|
|
160
|
+
relative = f"{variant_id}/model.gguf"
|
|
161
|
+
target = bundle.path(relative)
|
|
162
|
+
target.parent.mkdir(parents=True, exist_ok=True)
|
|
163
|
+
shutil.copyfile(source, target)
|
|
164
|
+
manifest = Manifest(
|
|
165
|
+
id=model_id or _bundle_id(info.metadata.get("name", source.stem)),
|
|
166
|
+
version="1.0.0",
|
|
167
|
+
task=Task.TEXT_GENERATION,
|
|
168
|
+
license=license,
|
|
169
|
+
name=info.metadata.get("name"),
|
|
170
|
+
source=f"file:{source.name}",
|
|
171
|
+
variants=[
|
|
172
|
+
Variant(
|
|
173
|
+
id=variant_id,
|
|
174
|
+
runtime=Runtime.LLAMACPP,
|
|
175
|
+
precision=quant,
|
|
176
|
+
file=bundle.add(relative),
|
|
177
|
+
min_ram_mb=estimate_min_ram_mb(info.size, context),
|
|
178
|
+
)
|
|
179
|
+
],
|
|
180
|
+
llm=LlmConfig(context_length=context),
|
|
181
|
+
)
|
|
182
|
+
return GgufImport(manifest, warnings)
|
|
183
|
+
|
|
184
|
+
|
|
185
|
+
def _card_license(info: Any) -> str | None:
|
|
186
|
+
card = getattr(info, "card_data", None)
|
|
187
|
+
value = getattr(card, "license", None) if card is not None else None
|
|
188
|
+
if isinstance(value, list):
|
|
189
|
+
value = value[0] if value else None
|
|
190
|
+
return str(value) if value else None
|
|
191
|
+
|
|
192
|
+
|
|
193
|
+
def _lfs_sha256(sibling: Any) -> str | None:
|
|
194
|
+
lfs = getattr(sibling, "lfs", None)
|
|
195
|
+
if lfs is None:
|
|
196
|
+
return None
|
|
197
|
+
return lfs.get("sha256") if isinstance(lfs, dict) else getattr(lfs, "sha256", None)
|