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.
Files changed (26) hide show
  1. belfort_ml-2026.10.8.dev0/.gitignore +11 -0
  2. belfort_ml-2026.10.8.dev0/PKG-INFO +9 -0
  3. belfort_ml-2026.10.8.dev0/README.md +228 -0
  4. belfort_ml-2026.10.8.dev0/belfort_ml/__init__.py +72 -0
  5. belfort_ml-2026.10.8.dev0/belfort_ml/_compile/__init__.py +64 -0
  6. belfort_ml-2026.10.8.dev0/belfort_ml/_compile/_annotations.py +299 -0
  7. belfort_ml-2026.10.8.dev0/belfort_ml/_compile/_export.py +204 -0
  8. belfort_ml-2026.10.8.dev0/belfort_ml/errors.py +186 -0
  9. belfort_ml-2026.10.8.dev0/belfort_ml/gateway_connector/__init__.py +0 -0
  10. belfort_ml-2026.10.8.dev0/belfort_ml/gateway_connector/_poll.py +39 -0
  11. belfort_ml-2026.10.8.dev0/belfort_ml/gateway_connector/_upload.py +65 -0
  12. belfort_ml-2026.10.8.dev0/belfort_ml/gateway_connector/client.py +77 -0
  13. belfort_ml-2026.10.8.dev0/belfort_ml/gateway_connector/deploys.py +376 -0
  14. belfort_ml-2026.10.8.dev0/belfort_ml/gateway_connector/models.py +303 -0
  15. belfort_ml-2026.10.8.dev0/belfort_ml/gateway_connector/transport.py +263 -0
  16. belfort_ml-2026.10.8.dev0/belfort_ml/py.typed +0 -0
  17. belfort_ml-2026.10.8.dev0/belfort_ml/types.py +26 -0
  18. belfort_ml-2026.10.8.dev0/pyproject.toml +97 -0
  19. belfort_ml-2026.10.8.dev0/tests/conftest.py +80 -0
  20. belfort_ml-2026.10.8.dev0/tests/test_deploys.py +88 -0
  21. belfort_ml-2026.10.8.dev0/tests/test_export.py +212 -0
  22. belfort_ml-2026.10.8.dev0/tests/test_export_e2e.py +136 -0
  23. belfort_ml-2026.10.8.dev0/tests/test_fetch_artifacts.py +139 -0
  24. belfort_ml-2026.10.8.dev0/tests/test_models.py +75 -0
  25. belfort_ml-2026.10.8.dev0/tests/test_transport.py +209 -0
  26. belfort_ml-2026.10.8.dev0/tests/test_upload.py +44 -0
@@ -0,0 +1,11 @@
1
+ __pycache__/
2
+ *.py[cod]
3
+ *.egg-info/
4
+ .mypy_cache/
5
+ .pytest_cache/
6
+ .ruff_cache/
7
+ .coverage
8
+ htmlcov/
9
+ .venv/
10
+ build/
11
+ dist/
@@ -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