belfort-ml 2026.10.8.dev0__tar.gz
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.
- belfort_ml-2026.10.8.dev0/.gitignore +11 -0
- belfort_ml-2026.10.8.dev0/PKG-INFO +9 -0
- belfort_ml-2026.10.8.dev0/README.md +228 -0
- belfort_ml-2026.10.8.dev0/belfort_ml/__init__.py +72 -0
- belfort_ml-2026.10.8.dev0/belfort_ml/_compile/__init__.py +64 -0
- belfort_ml-2026.10.8.dev0/belfort_ml/_compile/_annotations.py +299 -0
- belfort_ml-2026.10.8.dev0/belfort_ml/_compile/_export.py +204 -0
- belfort_ml-2026.10.8.dev0/belfort_ml/errors.py +186 -0
- belfort_ml-2026.10.8.dev0/belfort_ml/gateway_connector/__init__.py +0 -0
- belfort_ml-2026.10.8.dev0/belfort_ml/gateway_connector/_poll.py +39 -0
- belfort_ml-2026.10.8.dev0/belfort_ml/gateway_connector/_upload.py +65 -0
- belfort_ml-2026.10.8.dev0/belfort_ml/gateway_connector/client.py +77 -0
- belfort_ml-2026.10.8.dev0/belfort_ml/gateway_connector/deploys.py +376 -0
- belfort_ml-2026.10.8.dev0/belfort_ml/gateway_connector/models.py +303 -0
- belfort_ml-2026.10.8.dev0/belfort_ml/gateway_connector/transport.py +263 -0
- belfort_ml-2026.10.8.dev0/belfort_ml/py.typed +0 -0
- belfort_ml-2026.10.8.dev0/belfort_ml/types.py +26 -0
- belfort_ml-2026.10.8.dev0/pyproject.toml +97 -0
- belfort_ml-2026.10.8.dev0/tests/conftest.py +80 -0
- belfort_ml-2026.10.8.dev0/tests/test_deploys.py +88 -0
- belfort_ml-2026.10.8.dev0/tests/test_export.py +212 -0
- belfort_ml-2026.10.8.dev0/tests/test_export_e2e.py +136 -0
- belfort_ml-2026.10.8.dev0/tests/test_fetch_artifacts.py +139 -0
- belfort_ml-2026.10.8.dev0/tests/test_models.py +75 -0
- belfort_ml-2026.10.8.dev0/tests/test_transport.py +209 -0
- belfort_ml-2026.10.8.dev0/tests/test_upload.py +44 -0
|
@@ -0,0 +1,9 @@
|
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
|
+
Name: belfort-ml
|
|
3
|
+
Version: 2026.10.8.dev0
|
|
4
|
+
Summary: Belfort ML client for the Belfort CaaS gateway
|
|
5
|
+
Requires-Python: <3.13,>=3.12
|
|
6
|
+
Requires-Dist: belfort-ml-bucket==2026.10.8.dev0
|
|
7
|
+
Requires-Dist: belfort-torch-mlir==2026.6.25.dev0
|
|
8
|
+
Requires-Dist: httpx==0.28.1
|
|
9
|
+
Requires-Dist: torch==2.13.0
|
|
@@ -0,0 +1,228 @@
|
|
|
1
|
+
# Belfort ML Python Client
|
|
2
|
+
|
|
3
|
+
`belfort_ml`: an httpx client for the Belfort CaaS gateway. A model provider
|
|
4
|
+
uses it to compile a PyTorch model, register it, and run inference on it while
|
|
5
|
+
it stays encrypted.
|
|
6
|
+
|
|
7
|
+
## Install
|
|
8
|
+
|
|
9
|
+
Not published. It is a `uv` workspace member, so run `uv sync` from the repo
|
|
10
|
+
root.
|
|
11
|
+
|
|
12
|
+
## Configuration
|
|
13
|
+
|
|
14
|
+
```python
|
|
15
|
+
import belfort_ml
|
|
16
|
+
|
|
17
|
+
client = belfort_ml.MP_Client() # reads the environment
|
|
18
|
+
client = belfort_ml.MP_Client(api_key="blf_live_...") # or pass it in
|
|
19
|
+
```
|
|
20
|
+
|
|
21
|
+
| Argument | Environment | Default |
|
|
22
|
+
|---|---|---|
|
|
23
|
+
| `api_key` | `BELFORT_API_KEY` | required |
|
|
24
|
+
| `base_url` | `BELFORT_BASE_URL` | `https://eml.belfortlabs.com/api/v1` |
|
|
25
|
+
| `timeouts` | — | `Timeouts()` |
|
|
26
|
+
| `max_retries` | — | `3` |
|
|
27
|
+
|
|
28
|
+
Explicit arguments beat the environment. A missing key raises
|
|
29
|
+
`ConfigurationError` at construction. `Timeouts` is per operation class:
|
|
30
|
+
`standard` 30s, `upload_part` 300s, `inference` 120s, `poll` 10s.
|
|
31
|
+
|
|
32
|
+
## Quick start
|
|
33
|
+
|
|
34
|
+
```python
|
|
35
|
+
with belfort_ml.MP_Client() as client:
|
|
36
|
+
account = client.account.get()
|
|
37
|
+
print(account["mp_id"], account["credits"]["available"])
|
|
38
|
+
```
|
|
39
|
+
|
|
40
|
+
The client holds a connection pool for its lifetime, so use it as a context
|
|
41
|
+
manager or call `close()`.
|
|
42
|
+
|
|
43
|
+
`client.account` and `client.credits` return plain dicts. `client.models` and
|
|
44
|
+
`client.deploys` return typed objects — `Model`, `ClientKeySet`, `Deploy`,
|
|
45
|
+
`LoadedKeys`, `InferenceResult` — that carry their own methods, so
|
|
46
|
+
`model.deploy()` works without passing ids around.
|
|
47
|
+
|
|
48
|
+
## The gateway surface
|
|
49
|
+
|
|
50
|
+
| Method | Endpoint |
|
|
51
|
+
|---|---|
|
|
52
|
+
| `client.account.get()` | `GET /account` |
|
|
53
|
+
| `client.credits.get()` | `GET /credits` |
|
|
54
|
+
| `client.credits.transactions(cursor=, limit=, kind=)` | `GET /credits/transactions` |
|
|
55
|
+
| `client.models.create(artifacts, name=)` | `POST /models`, S3 PUTs, `POST /models/{m_id}/compile`, then polls |
|
|
56
|
+
| `client.models.get(id)` / `.list()` | `GET /models/{m_id}` / `GET /models` |
|
|
57
|
+
| `model.refresh()` / `.delete()` | `GET .../status` / `DELETE /models/{m_id}` |
|
|
58
|
+
| `model.fetch_artifacts(dir, bucket_url=)` | Reads the compiled artifacts from the bucket |
|
|
59
|
+
| `model.create_client(label, ttl_seconds=)` | `POST /models/{m_id}/client-tokens`: the token one end client runs on |
|
|
60
|
+
| `model.artifact_downloads()` | One presigned GET per artifact of the model's compile |
|
|
61
|
+
| `model.clients(status=)` / `model.delete_client(c_id)` | `GET /models/{m_id}/clients`, `DELETE /models/{m_id}/clients/{c_id}` |
|
|
62
|
+
| `model.deploy(key_set=, max_lifetime=)` | `POST /deploys` for the key set it serves, then polls |
|
|
63
|
+
| `client.deploys.list()` / `.get(id)` / `.stop_all()` | `GET /deploys`, `GET /deploys/{d_id}`, `DELETE` each running one |
|
|
64
|
+
| `deploy.refresh()` / `.extend(n)` / `.stop()` | `GET .../status` / `POST .../extend` / `DELETE /deploys/{d_id}` |
|
|
65
|
+
| `deploy.load_keys()` / `.unload_keys(k)` | `POST .../keys` for the deploy's key set / `DELETE .../keys/{k_id}` |
|
|
66
|
+
| `deploy.infer(loaded, bytes)` / `.infer_batch(...)` | `POST /deploys/{d_id}/infer` |
|
|
67
|
+
|
|
68
|
+
Nothing calls `GET /jobs/{job_id}`. Every `wait=True` path polls the resource's
|
|
69
|
+
own `/status` route: 2s backing off to 10s, then `belfort_ml.TimeoutError`. A ^C
|
|
70
|
+
while a deploy starts prints the deploy id and that the instance is still
|
|
71
|
+
billing, then re-raises.
|
|
72
|
+
|
|
73
|
+
`deploy(...)` is a context manager. Exit stops the instance whether the block
|
|
74
|
+
completed or raised, and a deploy someone else already stopped is not an error.
|
|
75
|
+
|
|
76
|
+
`client.credits.transactions` returns one page, newest first, keyset-paginated
|
|
77
|
+
on `tx_id`; follow `next_cursor` yourself. The models and deploys lists follow
|
|
78
|
+
cursors lazily.
|
|
79
|
+
|
|
80
|
+
Uploads go straight to S3 through presigned multipart URLs, without the Belfort
|
|
81
|
+
credential. The file is split into the presigned part count, with 4 parts in
|
|
82
|
+
flight, each retried on its own and streamed from disk.
|
|
83
|
+
|
|
84
|
+
## Compiling locally
|
|
85
|
+
|
|
86
|
+
```python
|
|
87
|
+
from belfort_ml._compile._export import load_calibration_data, load_pytorch_model
|
|
88
|
+
|
|
89
|
+
model = load_pytorch_model("smallnet.pt") # eval mode, CPU
|
|
90
|
+
calibration = load_calibration_data( # 64 examples by default
|
|
91
|
+
"calibration.npz", key="inputs", input_shape=(1, 3, 32, 32)
|
|
92
|
+
)
|
|
93
|
+
artifacts = belfort_ml.compile(
|
|
94
|
+
model, calibration[0], "smallnet", "./build", calibration_inputs=calibration
|
|
95
|
+
)
|
|
96
|
+
```
|
|
97
|
+
|
|
98
|
+
`compile` writes `build/smallnet_torch.mlir` and returns the `Artifacts` that
|
|
99
|
+
`client.models.create` takes: the file and the parameters it declares. No
|
|
100
|
+
network is involved.
|
|
101
|
+
|
|
102
|
+
`name` becomes the generated C++ code's name, so it must be a C++ identifier:
|
|
103
|
+
ASCII letters, digits and underscores, not starting with a digit, at most 200
|
|
104
|
+
characters, with no `__` and no leading `_` before an uppercase letter. It may
|
|
105
|
+
not be a C++ keyword, `std`, `heir`, `detail` or one of the native build's
|
|
106
|
+
macros (`PERSEUS_MODEL`, `PERSEUS_HEADER`, `PERSEUS_SERVER`, `DEFAULT_STREAM`)
|
|
107
|
+
or a standard-library macro (`NULL`, `EOF`, `errno`, `stdin`, `stdout`, `stderr`). Any other name raises `ValueError`
|
|
108
|
+
before export.
|
|
109
|
+
|
|
110
|
+
`example_input` fixes the shapes. The graph is specialized to it, so a deployed
|
|
111
|
+
model serves that input shape only. A model the frontend cannot lower raises
|
|
112
|
+
`CompileError` carrying the frontend's message.
|
|
113
|
+
|
|
114
|
+
Two annotations travel with the IR. Every input carries `{secret.secret}` unless
|
|
115
|
+
the model marks a subset with `belfort_ml.Secret`, as
|
|
116
|
+
[`examples/mlp/mlp.py`](../examples/mlp/mlp.py) does; and each op the compiler
|
|
117
|
+
must approximate carries the value range the calibration examples put through
|
|
118
|
+
it. Only those bounds reach the IR, never the examples.
|
|
119
|
+
|
|
120
|
+
The end client is its own component, [`belfort-ml-client/`](../belfort-ml-client).
|
|
121
|
+
|
|
122
|
+
## Reading the compiled artifacts
|
|
123
|
+
|
|
124
|
+
Perseus writes under `{mp_id}/{m_id}/compiles/{compile_id}/`. The gateway
|
|
125
|
+
reports a status and the chosen parameters; the files themselves are read from
|
|
126
|
+
the bucket:
|
|
127
|
+
|
|
128
|
+
```python
|
|
129
|
+
compiled = model.fetch_artifacts(
|
|
130
|
+
"./artifacts", bucket_url="http://localhost:9000/belfort-caas-models"
|
|
131
|
+
)
|
|
132
|
+
compiled.compile_id # the model's compile
|
|
133
|
+
compiled.params # what the compiler actually chose
|
|
134
|
+
compiled.entry_function # the exported function the C++ is named after
|
|
135
|
+
compiled.files # every downloaded file, tree preserved
|
|
136
|
+
```
|
|
137
|
+
|
|
138
|
+
| Argument | Environment | Default |
|
|
139
|
+
|---|---|---|
|
|
140
|
+
| `bucket_url` | `BELFORT_BUCKET_URL` | required |
|
|
141
|
+
| `access_key_id` | `BELFORT_BUCKET_ACCESS_KEY_ID` | boto3's own chain |
|
|
142
|
+
| `secret_access_key` | `BELFORT_BUCKET_SECRET_ACCESS_KEY` | boto3's own chain |
|
|
143
|
+
| `region` | `BELFORT_BUCKET_REGION` | `us-east-1` |
|
|
144
|
+
|
|
145
|
+
`bucket_url` names the endpoint and the bucket together, as
|
|
146
|
+
[`bucket/`](../bucket) describes; a URL carrying a key prefix raises
|
|
147
|
+
`ConfigurationError`.
|
|
148
|
+
|
|
149
|
+
It reads the model's compile, `model.active_compile_id`. Pass `store=` a
|
|
150
|
+
`belfort_bucket.Bucket` to reuse one reader across several models.
|
|
151
|
+
|
|
152
|
+
## The end client
|
|
153
|
+
|
|
154
|
+
`model.create_client(label)` opens a key set for one end client and returns a
|
|
155
|
+
`ClientToken`. Hand its `token` to the client on your own channel; the client
|
|
156
|
+
then talks to the gateway itself, with no bucket credential and no API key:
|
|
157
|
+
|
|
158
|
+
```bash
|
|
159
|
+
belfort-ml-client infer --token "$BELFORT_CLIENT_TOKEN" --input input.json
|
|
160
|
+
```
|
|
161
|
+
|
|
162
|
+
The client downloads and links its code, generates its key pair, uploads the
|
|
163
|
+
evaluation keys, and waits for the gateway to load them onto the deploy started
|
|
164
|
+
for its key set with `model.deploy(key_set=minted, max_lifetime=...)`. [`belfort-ml-client`](../belfort-ml-client) describes that side, and
|
|
165
|
+
[`e2e/`](../tests/e2e) runs both.
|
|
166
|
+
|
|
167
|
+
A key set travels as a directory, one object per evaluation key: `<n>.evk` and
|
|
168
|
+
one `<n>.s2d`.
|
|
169
|
+
|
|
170
|
+
## Errors
|
|
171
|
+
|
|
172
|
+
Everything descends from `BelfortMLError`, and local failures separate from
|
|
173
|
+
gateway failures.
|
|
174
|
+
|
|
175
|
+
| Raised locally | When |
|
|
176
|
+
|---|---|
|
|
177
|
+
| `ConfigurationError` | No API key, or a `base_url` that is not http(s) |
|
|
178
|
+
| `CompileError` | Local or gateway-side compile failure |
|
|
179
|
+
| `UploadError` | An S3 part failed after retries |
|
|
180
|
+
| `DownloadError` | Reading the bucket failed, or it holds no such compile |
|
|
181
|
+
| `KeyMaterialError` | A key directory holds something that may not be published |
|
|
182
|
+
| `TimeoutError` | Poll deadline exceeded |
|
|
183
|
+
| `ConnectionError` | Gateway unreachable after retries; httpx's error rides along as `__cause__` |
|
|
184
|
+
|
|
185
|
+
`APIError` covers anything the gateway returned, carrying `message`, `code`,
|
|
186
|
+
`detail`, `status_code` and `request_id`. Quote `request_id` when you report a
|
|
187
|
+
failure; it locates the request in the gateway logs.
|
|
188
|
+
|
|
189
|
+
| Status | Class | Extra |
|
|
190
|
+
|---|---|---|
|
|
191
|
+
| 400, 413 | `ValidationError` | |
|
|
192
|
+
| 401 | `AuthenticationError` | |
|
|
193
|
+
| 402 | `InsufficientCredits` | `.required`, `.available` |
|
|
194
|
+
| 403 | `PermissionError` | |
|
|
195
|
+
| 404 | `NotFoundError` | |
|
|
196
|
+
| 409 | `ConflictError` | `DeployLimitReached` for `deploy_limit_reached`, with `.limit`, `.running` |
|
|
197
|
+
| 429 | `RateLimitError` | `.retry_after` |
|
|
198
|
+
| 503 | `ServiceUnavailable` | |
|
|
199
|
+
|
|
200
|
+
A response that is not the gateway's `{"error": {...}}` envelope — a proxy page,
|
|
201
|
+
a crash — degrades to the raw text.
|
|
202
|
+
|
|
203
|
+
## Retries, idempotency and logging
|
|
204
|
+
|
|
205
|
+
Connection errors, timeouts, 429 and 500/502/504 are retried `max_retries`
|
|
206
|
+
times, with exponential backoff and full jitter capped at 30s. A 429's
|
|
207
|
+
`Retry-After` wins over the computed backoff. 4xx is never retried, and neither
|
|
208
|
+
is inference. A read or write timeout can come after the gateway acted on the
|
|
209
|
+
request, so after one only GET and HEAD are sent again.
|
|
210
|
+
|
|
211
|
+
Idempotent requests carry one `Idempotency-Key` per logical operation, reused
|
|
212
|
+
across that operation's retries. The gateway does not honour it yet.
|
|
213
|
+
|
|
214
|
+
Logging goes to the standard `belfort_ml` logger and is silent without a
|
|
215
|
+
handler. DEBUG lines carry method, path, status and request id only, never
|
|
216
|
+
headers or bodies.
|
|
217
|
+
|
|
218
|
+
## Examples and tests
|
|
219
|
+
|
|
220
|
+
[`examples/`](./examples/) covers the account, credits and error handling.
|
|
221
|
+
|
|
222
|
+
```bash
|
|
223
|
+
cd gateway && docker compose up -d minio # the artifact tests need it
|
|
224
|
+
uv run pytest # from belfort-ml/
|
|
225
|
+
```
|
|
226
|
+
|
|
227
|
+
CI runs this package from
|
|
228
|
+
[`.github/workflows/test-packages.yml`](../.github/workflows/test-packages.yml).
|
|
@@ -0,0 +1,72 @@
|
|
|
1
|
+
__version__ = "2026.10.8.dev0"
|
|
2
|
+
|
|
3
|
+
from belfort_bucket import CompiledArtifacts
|
|
4
|
+
|
|
5
|
+
from belfort_ml._compile import compile # noqa: A004
|
|
6
|
+
from belfort_ml._compile._annotations import Secret
|
|
7
|
+
from belfort_ml.errors import (
|
|
8
|
+
APIError,
|
|
9
|
+
AuthenticationError,
|
|
10
|
+
BelfortMLError,
|
|
11
|
+
CompileError,
|
|
12
|
+
ConfigurationError,
|
|
13
|
+
ConflictError,
|
|
14
|
+
ConnectionError,
|
|
15
|
+
DeployLimitReached,
|
|
16
|
+
DownloadError,
|
|
17
|
+
InsufficientCredits,
|
|
18
|
+
KeyMaterialError,
|
|
19
|
+
NotFoundError,
|
|
20
|
+
PermissionError,
|
|
21
|
+
RateLimitError,
|
|
22
|
+
ServiceUnavailable,
|
|
23
|
+
TimeoutError,
|
|
24
|
+
UploadError,
|
|
25
|
+
ValidationError,
|
|
26
|
+
)
|
|
27
|
+
from belfort_ml.gateway_connector.client import MP_Client
|
|
28
|
+
from belfort_ml.gateway_connector.deploys import (
|
|
29
|
+
Deploy,
|
|
30
|
+
InferenceResult,
|
|
31
|
+
LoadedKeys,
|
|
32
|
+
)
|
|
33
|
+
from belfort_ml.gateway_connector.models import (
|
|
34
|
+
ClientKeySet,
|
|
35
|
+
ClientToken,
|
|
36
|
+
Model,
|
|
37
|
+
)
|
|
38
|
+
from belfort_ml.types import Artifacts, CompileParams
|
|
39
|
+
|
|
40
|
+
__all__ = [
|
|
41
|
+
"APIError",
|
|
42
|
+
"Artifacts",
|
|
43
|
+
"AuthenticationError",
|
|
44
|
+
"MP_Client",
|
|
45
|
+
"ClientKeySet",
|
|
46
|
+
"CompileError",
|
|
47
|
+
"CompileParams",
|
|
48
|
+
"CompiledArtifacts",
|
|
49
|
+
"Deploy",
|
|
50
|
+
"DeployLimitReached",
|
|
51
|
+
"ClientToken",
|
|
52
|
+
"InferenceResult",
|
|
53
|
+
"KeyMaterialError",
|
|
54
|
+
"LoadedKeys",
|
|
55
|
+
"Model",
|
|
56
|
+
"ConfigurationError",
|
|
57
|
+
"ConflictError",
|
|
58
|
+
"ConnectionError",
|
|
59
|
+
"DownloadError",
|
|
60
|
+
"InsufficientCredits",
|
|
61
|
+
"BelfortMLError",
|
|
62
|
+
"NotFoundError",
|
|
63
|
+
"PermissionError",
|
|
64
|
+
"RateLimitError",
|
|
65
|
+
"Secret",
|
|
66
|
+
"ServiceUnavailable",
|
|
67
|
+
"TimeoutError",
|
|
68
|
+
"UploadError",
|
|
69
|
+
"ValidationError",
|
|
70
|
+
"__version__",
|
|
71
|
+
"compile",
|
|
72
|
+
]
|
|
@@ -0,0 +1,64 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import re
|
|
4
|
+
from pathlib import Path
|
|
5
|
+
from typing import Any
|
|
6
|
+
|
|
7
|
+
from belfort_ml._compile._export import SampleInput, export_torch_mlir
|
|
8
|
+
from belfort_ml.types import Artifacts
|
|
9
|
+
|
|
10
|
+
# `name` names the C++ Perseus generates (`heir::generated::<name>`) and its
|
|
11
|
+
# artifacts, so it must be a C++ identifier that is neither a keyword, a
|
|
12
|
+
# namespace that code refers to, nor a macro the native build or the standard
|
|
13
|
+
# headers define, and short enough that `<name>_cheddar.mlir` fits a 255-byte
|
|
14
|
+
# file name.
|
|
15
|
+
_NAME = re.compile(r"[A-Za-z_][A-Za-z0-9_]{0,199}")
|
|
16
|
+
_RESERVED_NAMES = frozenset(
|
|
17
|
+
"""
|
|
18
|
+
alignas alignof and and_eq asm auto bitand bitor bool break case catch char
|
|
19
|
+
char8_t char16_t char32_t class compl concept const consteval constexpr
|
|
20
|
+
constinit const_cast continue co_await co_return co_yield decltype default
|
|
21
|
+
delete do double dynamic_cast else enum explicit export extern false float
|
|
22
|
+
for friend goto if inline int long mutable namespace new noexcept not not_eq
|
|
23
|
+
nullptr operator or or_eq private protected public register reinterpret_cast
|
|
24
|
+
requires return short signed sizeof static static_assert static_cast struct
|
|
25
|
+
switch template this thread_local throw true try typedef typeid typename
|
|
26
|
+
union unsigned using virtual void volatile wchar_t while xor xor_eq
|
|
27
|
+
detail heir std
|
|
28
|
+
DEFAULT_STREAM PERSEUS_HEADER PERSEUS_MODEL PERSEUS_SERVER
|
|
29
|
+
EOF NULL errno stderr stdin stdout
|
|
30
|
+
""".split()
|
|
31
|
+
)
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
def _validate_name(name: str) -> None:
|
|
35
|
+
if not _NAME.fullmatch(name):
|
|
36
|
+
raise ValueError(
|
|
37
|
+
f"model name {name!r} may use only ASCII letters, digits and "
|
|
38
|
+
"underscores, may not start with a digit, and may be at most 200 "
|
|
39
|
+
"characters long"
|
|
40
|
+
)
|
|
41
|
+
# C++ reserves these for the implementation (`__FILE__`, `_Bool`).
|
|
42
|
+
if name in _RESERVED_NAMES or "__" in name or re.match(r"_[A-Z]", name):
|
|
43
|
+
raise ValueError(
|
|
44
|
+
f"model name {name!r} is reserved in the C++ the compiler generates"
|
|
45
|
+
)
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
def compile( # noqa: A001
|
|
49
|
+
model: Any, # torch.nn.Module
|
|
50
|
+
example_input: SampleInput,
|
|
51
|
+
name: str,
|
|
52
|
+
output_dir: str | Path,
|
|
53
|
+
*,
|
|
54
|
+
calibration_inputs: list[SampleInput] | None = None,
|
|
55
|
+
) -> Artifacts:
|
|
56
|
+
_validate_name(name)
|
|
57
|
+
mlir_path = export_torch_mlir(
|
|
58
|
+
model,
|
|
59
|
+
example_input,
|
|
60
|
+
name,
|
|
61
|
+
output_dir,
|
|
62
|
+
calibration_inputs=calibration_inputs,
|
|
63
|
+
)
|
|
64
|
+
return Artifacts(mlir_path=mlir_path)
|
|
@@ -0,0 +1,299 @@
|
|
|
1
|
+
"""Secret-input markers and range calibration, applied inside `export_and_import`'s `annotate=` callback."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import math
|
|
6
|
+
from collections.abc import Iterable
|
|
7
|
+
from typing import Any, get_type_hints
|
|
8
|
+
|
|
9
|
+
import torch
|
|
10
|
+
import torch.nn as nn
|
|
11
|
+
from torch.export import ExportedProgram
|
|
12
|
+
from torch.export.graph_signature import InputKind
|
|
13
|
+
from torch_mlir.extras import fx_importer # type: ignore[import-untyped]
|
|
14
|
+
from torch_mlir.extras.annotate import annotate_arg # type: ignore[import-untyped]
|
|
15
|
+
|
|
16
|
+
from belfort_ml.errors import CompileError
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
class Secret:
|
|
20
|
+
"""PEP-593 marker: the annotated forward argument is a ciphertext input."""
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
SECRET_ATTR = {"secret.secret": True}
|
|
24
|
+
|
|
25
|
+
# Tail probability dropped at each end of a calibrated range; 0 gives the exact min/max.
|
|
26
|
+
DEFAULT_LOWER_QUANTILE = 0.01
|
|
27
|
+
|
|
28
|
+
# The calibrated half-range is multiplied by this.
|
|
29
|
+
DEFAULT_DOMAIN_MARGIN = 2.0
|
|
30
|
+
|
|
31
|
+
# Values sampled per recorded node.
|
|
32
|
+
CAPACITY = 1 << 14
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def apply_forward_annotations(prog: ExportedProgram, model: nn.Module) -> int:
|
|
36
|
+
"""Mark the forward arguments annotated with `Secret`; return how many."""
|
|
37
|
+
try:
|
|
38
|
+
hints = get_type_hints(model.forward, include_extras=True)
|
|
39
|
+
except (NameError, TypeError) as exc:
|
|
40
|
+
raise CompileError(
|
|
41
|
+
f"could not read the annotations of {type(model).__name__}.forward: "
|
|
42
|
+
f"{exc}. Make every annotation resolvable from the module that "
|
|
43
|
+
f"defines the model, so the `Secret` markers can be read."
|
|
44
|
+
) from exc
|
|
45
|
+
marked = 0
|
|
46
|
+
for name, hint in hints.items():
|
|
47
|
+
if name != "return" and Secret in getattr(hint, "__metadata__", ()):
|
|
48
|
+
annotate_arg(prog, name, dict(SECRET_ATTR))
|
|
49
|
+
marked += 1
|
|
50
|
+
return marked
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
def annotate_all_inputs_secret(prog: ExportedProgram) -> int:
|
|
54
|
+
"""Mark every user input as a ciphertext input; return how many."""
|
|
55
|
+
n_inputs = sum(
|
|
56
|
+
1
|
|
57
|
+
for spec in prog.graph_signature.input_specs
|
|
58
|
+
if spec.kind == InputKind.USER_INPUT
|
|
59
|
+
)
|
|
60
|
+
for index in range(n_inputs):
|
|
61
|
+
annotate_arg(prog, index, dict(SECRET_ATTR))
|
|
62
|
+
return n_inputs
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
class ValueSample:
|
|
66
|
+
"""Reservoir sample (vectorized Algorithm R) of a stream of tensors, plus the exact min and max."""
|
|
67
|
+
|
|
68
|
+
def __init__(self, generator: torch.Generator) -> None:
|
|
69
|
+
self._generator = generator
|
|
70
|
+
self._reservoir = torch.empty(0, dtype=torch.float32)
|
|
71
|
+
self.count = 0
|
|
72
|
+
self._min = math.inf
|
|
73
|
+
self._max = -math.inf
|
|
74
|
+
|
|
75
|
+
def update(self, t: torch.Tensor) -> None:
|
|
76
|
+
values = t.detach().reshape(-1).to(torch.float32)
|
|
77
|
+
if values.numel() == 0:
|
|
78
|
+
return
|
|
79
|
+
self._min = min(self._min, float(values.min()))
|
|
80
|
+
self._max = max(self._max, float(values.max()))
|
|
81
|
+
|
|
82
|
+
free = CAPACITY - self._reservoir.numel()
|
|
83
|
+
if free > 0:
|
|
84
|
+
take = min(free, values.numel())
|
|
85
|
+
self._reservoir = torch.cat([self._reservoir, values[:take]])
|
|
86
|
+
self.count += take
|
|
87
|
+
values = values[take:]
|
|
88
|
+
if values.numel() > 0:
|
|
89
|
+
# Stream index i replaces a random slot with probability CAPACITY/(i+1).
|
|
90
|
+
index = torch.arange(self.count, self.count + values.numel())
|
|
91
|
+
accept = torch.rand(
|
|
92
|
+
values.numel(), generator=self._generator
|
|
93
|
+
) < CAPACITY / (index + 1).to(torch.float32)
|
|
94
|
+
n_accept = int(accept.sum())
|
|
95
|
+
if n_accept > 0:
|
|
96
|
+
slots = torch.randint(
|
|
97
|
+
0, CAPACITY, (n_accept,), generator=self._generator
|
|
98
|
+
)
|
|
99
|
+
self._reservoir[slots] = values[accept]
|
|
100
|
+
self.count += values.numel()
|
|
101
|
+
|
|
102
|
+
def quantile(self, q: float) -> float:
|
|
103
|
+
if q <= 0.0:
|
|
104
|
+
return self._min
|
|
105
|
+
if q >= 1.0:
|
|
106
|
+
return self._max
|
|
107
|
+
return float(torch.quantile(self._reservoir, q))
|
|
108
|
+
|
|
109
|
+
|
|
110
|
+
# Aten ops that lower to an elementwise `linalg.generic` and that the compiler
|
|
111
|
+
# approximates by a polynomial; the CKKS-exact add/sub/mul/neg are absent.
|
|
112
|
+
_ELEMENTWISE_OP_NAMES = (
|
|
113
|
+
"relu", "relu6", "leaky_relu", "prelu", "gelu", "silu", "mish",
|
|
114
|
+
"elu", "celu", "selu", "sigmoid", "hardsigmoid", "hardswish",
|
|
115
|
+
"hardtanh", "hardshrink", "softshrink", "softplus", "threshold",
|
|
116
|
+
"tanh", "erf", "logit",
|
|
117
|
+
"exp", "exp2", "expm1", "log", "log1p", "log2", "log10",
|
|
118
|
+
"sqrt", "rsqrt", "reciprocal", "pow",
|
|
119
|
+
"sin", "cos", "tan", "asin", "acos", "atan", "atan2",
|
|
120
|
+
"sinh", "cosh", "asinh", "acosh", "atanh",
|
|
121
|
+
"div", "remainder", "fmod",
|
|
122
|
+
"abs", "sign", "floor", "ceil", "round", "trunc",
|
|
123
|
+
"maximum", "minimum", "clamp", "clamp_min", "clamp_max",
|
|
124
|
+
"where", "lerp",
|
|
125
|
+
) # fmt: skip
|
|
126
|
+
|
|
127
|
+
# Ops whose domain is not all of R: widening could cross a singularity, so they
|
|
128
|
+
# keep their observed range.
|
|
129
|
+
_RESTRICTED_DOMAIN_OP_NAMES = (
|
|
130
|
+
"sqrt", "rsqrt", "reciprocal", "pow",
|
|
131
|
+
"log", "log1p", "log2", "log10", "logit",
|
|
132
|
+
"div", "remainder", "fmod",
|
|
133
|
+
"asin", "acos", "atanh", "acosh",
|
|
134
|
+
) # fmt: skip
|
|
135
|
+
|
|
136
|
+
|
|
137
|
+
def _aten_ops(names: tuple[str, ...]) -> frozenset[Any]:
|
|
138
|
+
return frozenset(
|
|
139
|
+
packet
|
|
140
|
+
for name in names
|
|
141
|
+
if (packet := getattr(torch.ops.aten, name, None)) is not None
|
|
142
|
+
)
|
|
143
|
+
|
|
144
|
+
|
|
145
|
+
ELEMENTWISE_ATEN_OPS = _aten_ops(_ELEMENTWISE_OP_NAMES)
|
|
146
|
+
RESTRICTED_DOMAIN_ATEN_OPS = _aten_ops(_RESTRICTED_DOMAIN_OP_NAMES)
|
|
147
|
+
|
|
148
|
+
|
|
149
|
+
def _targets(node: torch.fx.Node, ops: frozenset[Any]) -> bool:
|
|
150
|
+
return (
|
|
151
|
+
node.op == "call_function"
|
|
152
|
+
and getattr(node.target, "overloadpacket", None) in ops
|
|
153
|
+
)
|
|
154
|
+
|
|
155
|
+
|
|
156
|
+
def lowers_to_elementwise(node: torch.fx.Node) -> bool:
|
|
157
|
+
return _targets(node, ELEMENTWISE_ATEN_OPS)
|
|
158
|
+
|
|
159
|
+
|
|
160
|
+
class RangeRecorder(torch.fx.Interpreter):
|
|
161
|
+
"""Interprets an exported program, sampling the values that feed each
|
|
162
|
+
elementwise-lowered op: user inputs and `call_function` results, never
|
|
163
|
+
parameters, buffers or constants.
|
|
164
|
+
"""
|
|
165
|
+
|
|
166
|
+
def __init__(self, prog: ExportedProgram) -> None:
|
|
167
|
+
super().__init__(prog.graph_module)
|
|
168
|
+
self._prog = prog
|
|
169
|
+
# Seeded, so calibration is repeatable.
|
|
170
|
+
self._generator = torch.Generator().manual_seed(0)
|
|
171
|
+
self._user_inputs = {
|
|
172
|
+
spec.arg.name
|
|
173
|
+
for spec in prog.graph_signature.input_specs
|
|
174
|
+
if spec.kind == InputKind.USER_INPUT
|
|
175
|
+
}
|
|
176
|
+
# Parameters, buffers and constants: lifted to inputs the caller does not supply.
|
|
177
|
+
self._lifted = {**prog.state_dict, **prog.constants}
|
|
178
|
+
# Producers feeding an elementwise-lowered op.
|
|
179
|
+
self._monitored = {
|
|
180
|
+
src.name
|
|
181
|
+
for node in prog.graph.nodes
|
|
182
|
+
if lowers_to_elementwise(node)
|
|
183
|
+
for src in node.all_input_nodes
|
|
184
|
+
}
|
|
185
|
+
self.samples: dict[str, ValueSample] = {}
|
|
186
|
+
|
|
187
|
+
def run_node(self, n: torch.fx.Node) -> Any:
|
|
188
|
+
result = super().run_node(n)
|
|
189
|
+
recorded = n.op == "call_function" or (
|
|
190
|
+
n.op == "placeholder" and n.name in self._user_inputs
|
|
191
|
+
)
|
|
192
|
+
if (
|
|
193
|
+
n.name in self._monitored
|
|
194
|
+
and recorded
|
|
195
|
+
and isinstance(result, torch.Tensor)
|
|
196
|
+
and result.is_floating_point()
|
|
197
|
+
):
|
|
198
|
+
sample = self.samples.setdefault(n.name, ValueSample(self._generator))
|
|
199
|
+
sample.update(result)
|
|
200
|
+
return result
|
|
201
|
+
|
|
202
|
+
def run_calibration(self, batch: torch.Tensor | tuple[Any, ...]) -> None:
|
|
203
|
+
"""Run one calibration example, shaped like the export sample, in forward-argument order."""
|
|
204
|
+
user_inputs = iter((batch,) if isinstance(batch, torch.Tensor) else batch)
|
|
205
|
+
args = [
|
|
206
|
+
next(user_inputs)
|
|
207
|
+
if spec.kind == InputKind.USER_INPUT
|
|
208
|
+
else self._lifted[spec.target]
|
|
209
|
+
for spec in self._prog.graph_signature.input_specs
|
|
210
|
+
]
|
|
211
|
+
with torch.no_grad():
|
|
212
|
+
self.run(*args)
|
|
213
|
+
|
|
214
|
+
def bounds(self, q: float) -> dict[str, tuple[float, float]]:
|
|
215
|
+
"""Per-op input domains: for each elementwise-lowered node, the union of
|
|
216
|
+
its producers' [q, 1-q] output ranges. Nodes with no recorded producer —
|
|
217
|
+
one fed only by weights, say — are omitted.
|
|
218
|
+
"""
|
|
219
|
+
if not 0 <= q < 0.5:
|
|
220
|
+
raise ValueError(f"tail probability q must satisfy 0 <= q < 0.5, got {q}")
|
|
221
|
+
out: dict[str, tuple[float, float]] = {}
|
|
222
|
+
for node in self._prog.graph.nodes:
|
|
223
|
+
if not lowers_to_elementwise(node):
|
|
224
|
+
continue
|
|
225
|
+
lo, hi = math.inf, -math.inf
|
|
226
|
+
for src in node.all_input_nodes:
|
|
227
|
+
sample = self.samples.get(src.name)
|
|
228
|
+
if sample is not None and sample.count > 0:
|
|
229
|
+
lo = min(lo, sample.quantile(q))
|
|
230
|
+
hi = max(hi, sample.quantile(1 - q))
|
|
231
|
+
if lo <= hi:
|
|
232
|
+
out[node.name] = (lo, hi)
|
|
233
|
+
return out
|
|
234
|
+
|
|
235
|
+
|
|
236
|
+
def record_ranges(
|
|
237
|
+
prog: ExportedProgram,
|
|
238
|
+
calibration_data: Iterable[torch.Tensor | tuple[Any, ...]],
|
|
239
|
+
*,
|
|
240
|
+
q: float = DEFAULT_LOWER_QUANTILE,
|
|
241
|
+
margin: float = DEFAULT_DOMAIN_MARGIN,
|
|
242
|
+
) -> dict[str, tuple[float, float]]:
|
|
243
|
+
"""`{fx_node_name: (domain_lo, domain_hi)}` for `apply_ranges`.
|
|
244
|
+
|
|
245
|
+
The domain is the [q, 1-q] range of the values the op was evaluated at,
|
|
246
|
+
widened about its centre by `margin`; restricted-domain ops keep the raw range.
|
|
247
|
+
"""
|
|
248
|
+
recorder = RangeRecorder(prog)
|
|
249
|
+
for batch in calibration_data:
|
|
250
|
+
recorder.run_calibration(batch)
|
|
251
|
+
|
|
252
|
+
nodes = {node.name: node for node in prog.graph.nodes}
|
|
253
|
+
out: dict[str, tuple[float, float]] = {}
|
|
254
|
+
for name, (lo, hi) in recorder.bounds(q).items():
|
|
255
|
+
if _targets(nodes[name], RESTRICTED_DOMAIN_ATEN_OPS):
|
|
256
|
+
out[name] = (lo, hi)
|
|
257
|
+
else:
|
|
258
|
+
centre, half = (lo + hi) / 2, (hi - lo) / 2 * margin
|
|
259
|
+
out[name] = (centre - half, centre + half)
|
|
260
|
+
return out
|
|
261
|
+
|
|
262
|
+
|
|
263
|
+
def apply_ranges(prog: ExportedProgram, bounds: dict[str, tuple[float, float]]) -> int:
|
|
264
|
+
"""Attach `{domain_lower, domain_upper}` to the named nodes; return how many."""
|
|
265
|
+
nodes = {node.name: node for node in prog.graph.nodes}
|
|
266
|
+
hits = 0
|
|
267
|
+
for name, (lo, hi) in bounds.items():
|
|
268
|
+
node = nodes.get(name)
|
|
269
|
+
if node is None:
|
|
270
|
+
continue
|
|
271
|
+
# Rounded, so the emitted MLIR is reproducible across machines.
|
|
272
|
+
node.meta.setdefault(fx_importer.MLIR_OP_ATTRS_META_KEY, {}).update(
|
|
273
|
+
{"domain_lower": round(float(lo), 4), "domain_upper": round(float(hi), 4)}
|
|
274
|
+
)
|
|
275
|
+
hits += 1
|
|
276
|
+
return hits
|
|
277
|
+
|
|
278
|
+
|
|
279
|
+
def freeze_buffers(prog: ExportedProgram) -> int:
|
|
280
|
+
"""Inline every unmutated buffer as a literal; return how many."""
|
|
281
|
+
signature = prog.graph_signature
|
|
282
|
+
mutated = set(signature.buffers_to_mutate.values())
|
|
283
|
+
# A `persistent=False` buffer lives in `constants`, not `state_dict`.
|
|
284
|
+
lifted = {**prog.state_dict, **prog.constants}
|
|
285
|
+
replacements = {
|
|
286
|
+
input_name: lifted[state_name]
|
|
287
|
+
for input_name, state_name in signature.inputs_to_buffers.items()
|
|
288
|
+
if state_name not in mutated
|
|
289
|
+
}
|
|
290
|
+
|
|
291
|
+
frozen = 0
|
|
292
|
+
for node in list(prog.graph.nodes):
|
|
293
|
+
if node.op != "placeholder" or node.name not in replacements:
|
|
294
|
+
continue
|
|
295
|
+
# The importer emits a substituted tensor as a `torch.vtensor.literal`.
|
|
296
|
+
node.replace_all_uses_with(replacements[node.name])
|
|
297
|
+
prog.graph.erase_node(node)
|
|
298
|
+
frozen += 1
|
|
299
|
+
return frozen
|