changh95 commited on
Commit
b8d0ee3
·
verified ·
1 Parent(s): b15da4f

Python API: warm-up so the first real call is fast (warmup_variants, model.warmup()), quiet logs, install extras

Browse files

from_pretrained() now warms up the common per-request variants with dummy inputs (device and host), so the first real call is within about 10% of later calls (independently verified in fresh processes; outputs bit-identical; steady-state latency unchanged). Adds model.warmup(...) for other variants, verbose=False by default, and documents the [server,test] extras and repo-root paths. tt-model.yaml and SERVING.md are unchanged.

PYTHON.md CHANGED
@@ -10,13 +10,36 @@ You need an environment that already has `ttnn` (tt-metal `v0.78.0-dev20260820`,
10
  Then install the `code/` directory of this repo:
11
 
12
  ```bash
13
- pip install -e code/ # or: pip install code/
14
- pip install -e "code/[pose]" # also installs OpenCV for return_pose=True
 
 
15
  ```
16
 
17
  The package declares only its own dependencies (numpy<2, torch, pillow, safetensors, huggingface_hub).
18
  It does not install ttnn. You do not need `sys.path` changes or a specific working directory.
19
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
20
  ## Quickstart
21
 
22
  ```python
@@ -29,8 +52,12 @@ with Mast3rP150.from_pretrained(device_id=0) as model:
29
  out.save_ply("pointcloud.ply")
30
  ```
31
 
32
- `code/examples/quickstart.py` runs this snippet on the demo pair. It writes `pointmaps.npz`, `pointcloud.ply`,
33
- `depth.png` and `summary.json`.
 
 
 
 
34
 
35
  ## `Mast3rP150.from_pretrained(...)`
36
 
@@ -41,10 +68,13 @@ with Mast3rP150.from_pretrained(device_id=0) as model:
41
  | `dispatch` | `"auto"` | `auto`: if the tt-metal tree has the ETH-dispatch patch, use ETH dispatch and a 12x10 grid (the published configuration). If not, use the default Tensix dispatch and show a warning. `eth` or `worker` forces one mode. Env: `MAST3R_DISPATCH`. |
42
  | `repo_id`, `revision`, `weights_dir` | `naver/DUSt3R_ViTLarge_BaseDecoder_512_dpt` at `61c57447` | The source of `model.safetensors`: the HF cache (downloaded on first use), or a local directory. Env: `HF_MODEL`, `TT_WEIGHTS_REVISION`, `MAST3R_WEIGHTS_DIR`. |
43
  | `preprocess` | `"pad"` | Makes each view square. `pad` adds gray padding. `crop` cuts a centre square. Env: `MAST3R_PREPROC`. |
44
- | `warmup` | `1` | Number of warm-up forwards. The first one compiles the kernels and captures the metal trace. |
45
- | `capture_pose` | `True` | Also captures the symmetric graph that `return_pose=True` uses. |
 
 
46
 
47
- Startup takes about 65 s with an empty kernel cache and about 7 s with a warm cache (weights in the HF cache).
 
48
  The device-graph knobs (`TT_FUSED`, `MAST3R_OPT`, `MAST3R_CQS`, ...) come from the environment, as in the server.
49
  Their defaults are the published configuration.
50
  You can open only one model at a time in a process. `model.info` shows the dispatch mode, the grid and the weights path.
@@ -106,6 +136,109 @@ output = model.inference(pairs) # DUSt3R inference() o
106
  `inference` returns `{'view1', 'view2', 'pred1', 'pred2', 'loss'}` with CPU torch tensors, in DUSt3R's layout.
107
  All views are 512x512 squares. Upstream DUSt3R keeps the aspect ratio, but this port does not.
108
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
109
  ## Lifetime
110
 
111
  Use `with Mast3rP150.from_pretrained() as model:` or call `model.close()`.
 
10
  Then install the `code/` directory of this repo:
11
 
12
  ```bash
13
+ pip install -e code/ # or: pip install code/
14
+ pip install -e "code/[pose]" # also installs OpenCV for return_pose=True
15
+ pip install -e "code/[server]" # also installs FastAPI and uvicorn for the HTTP server (SERVING.md)
16
+ pip install -e "code/[pose,server,test]" # everything the host tests and the device tests need
17
  ```
18
 
19
  The package declares only its own dependencies (numpy<2, torch, pillow, safetensors, huggingface_hub).
20
  It does not install ttnn. You do not need `sys.path` changes or a specific working directory.
21
 
22
+ | extra | installs | you need it for |
23
+ |---|---|---|
24
+ | `pose` | `opencv-python-headless` | `return_pose=True` (PnP-RANSAC) |
25
+ | `server` | `fastapi`, `uvicorn`, `pydantic>=2` | the HTTP server `models/server/app.py`, and the host tests that compare with the server |
26
+ | `test` | `pytest`, `httpx` | all tests (`httpx` is for FastAPI's `TestClient`) |
27
+ | `viz` | `matplotlib` | optional figures |
28
+
29
+ The host tests do not open a chip and do not need ttnn. Run only the host tests:
30
+
31
+ ```bash
32
+ cd code
33
+ TT_VISIBLE_DEVICES=none python -m pytest -q models/tests/test_api_host.py models/tests/test_server_host.py models/tests/test_fused_host.py
34
+ ```
35
+
36
+ The device tests need a chip, ttnn and the weights:
37
+
38
+ ```bash
39
+ cd code
40
+ timeout -s INT 1200 python -m pytest -q -s models/tests/test_api_device.py models/tests/test_warmup_device.py
41
+ ```
42
+
43
  ## Quickstart
44
 
45
  ```python
 
52
  out.save_ply("pointcloud.ply")
53
  ```
54
 
55
+ The paths `media/source_1.png` and `media/source_2.png` are relative to the repo root.
56
+ Run the snippet from the repo root, or give absolute paths.
57
+
58
+ `code/examples/quickstart.py` runs this snippet on the demo pair. It finds the demo images relative to its own
59
+ file, so you can run it from any directory. It writes `pointmaps.npz`, `pointcloud.ply`, `depth.png` and
60
+ `summary.json` into `--out-dir` (default `quickstart_out`, relative to the current directory).
61
 
62
  ## `Mast3rP150.from_pretrained(...)`
63
 
 
68
  | `dispatch` | `"auto"` | `auto`: if the tt-metal tree has the ETH-dispatch patch, use ETH dispatch and a 12x10 grid (the published configuration). If not, use the default Tensix dispatch and show a warning. `eth` or `worker` forces one mode. Env: `MAST3R_DISPATCH`. |
69
  | `repo_id`, `revision`, `weights_dir` | `naver/DUSt3R_ViTLarge_BaseDecoder_512_dpt` at `61c57447` | The source of `model.safetensors`: the HF cache (downloaded on first use), or a local directory. Env: `HF_MODEL`, `TT_WEIGHTS_REVISION`, `MAST3R_WEIGHTS_DIR`. |
70
  | `preprocess` | `"pad"` | Makes each view square. `pad` adds gray padding. `crop` cuts a centre square. Env: `MAST3R_PREPROC`. |
71
+ | `warmup` | `1` | Number of warm-up forwards on a zero pair. The first one compiles the kernels and captures the metal trace. |
72
+ | `warmup_variants` | `("pair", "pose", "batch")` | The call variants to prepare before `from_pretrained` returns, so that the first real call of each variant is fast. See [Warm-up](#warm-up). |
73
+ | `capture_pose` | `None` | Legacy switch. `False` removes `"pose"` from `warmup_variants` (no symmetric graph). `True` adds it. |
74
+ | `verbose` | `False` | `False`: tt-metal prints only errors and ttnn prints only warnings and errors. `True`: the full tt-metal log and this package's INFO log. Env: `MAST3R_VERBOSE=1`. See [Console output](#console-output). |
75
 
76
+ Startup takes about 67 s with an empty kernel cache and about 9 s with a warm cache (weights in the HF cache).
77
+ Of the 9 s, the default warm-up takes about 7 s (device graphs about 6 s, host about 1 s).
78
  The device-graph knobs (`TT_FUSED`, `MAST3R_OPT`, `MAST3R_CQS`, ...) come from the environment, as in the server.
79
  Their defaults are the published configuration.
80
  You can open only one model at a time in a process. `model.info` shows the dispatch mode, the grid and the weights path.
 
136
  `inference` returns `{'view1', 'view2', 'pred1', 'pred2', 'loss'}` with CPU torch tensors, in DUSt3R's layout.
137
  All views are 512x512 squares. Upstream DUSt3R keeps the aspect ratio, but this port does not.
138
 
139
+ ## Warm-up
140
+
141
+ `from_pretrained` prepares the call variants in `warmup_variants` before it returns.
142
+ Thus the first real call of each prepared variant is as fast as the later calls.
143
+ The device graph is fixed (512x512, one pair), so the variants are few:
144
+
145
+ | variant | call it prepares | device | host first-use costs it pays |
146
+ |---|---|---|---|
147
+ | `"pair"` | `model(v1, v2)`, `forward_raw` | pair graph: kernel compile, trace capture, persistent input | PNG and JPEG decode, pad resize, numpy-array input, host thread pool, readback and activation buffers, `point_cloud` |
148
+ | `"pose"` | `model(v1, v2, return_pose=True)` | symmetric graph (trace capture) | OpenCV import (about 65 ms), focal vote, one PnP-RANSAC |
149
+ | `"batch"` | `predict_pairs`, `inference`, `load_images` / `make_pairs` | none (uses the pair graph) | the pipelined submit / collect path with 2 and 4 pairs |
150
+ | `"crop"` | `model(..., preprocess="crop")` | none (uses the pair graph) | the crop resize |
151
+
152
+ The default is `("pair", "pose", "batch")`. `"all"` adds `"crop"`.
153
+ A variant can also be a dict of call options: `{"return_pose": True, "preprocess": "crop"}` (keys `return_pose`, `batch`, `preprocess`).
154
+
155
+ ```python
156
+ model = Mast3rP150.from_pretrained(warmup_variants="all") # also prepare preprocess="crop"
157
+ model = Mast3rP150.from_pretrained(warmup_variants=["pair"]) # faster start-up; no symmetric graph
158
+ model.warmup(return_pose=True) # prepare a variant later; returns {"pose": ms}
159
+ model.warmup("pose") # again: {} (idempotent, no device work)
160
+ model.warmed_variants # ['pair', 'pose']
161
+ model.info["warmup_ms"] # time of each variant warm-up
162
+ ```
163
+
164
+ How the warm-up works:
165
+
166
+ - The device graph of each variant is compiled and captured on a zero pair, as the server does.
167
+ - Then each variant makes 2 complete calls through the public API with realistic dummy images.
168
+ The images are a 640x480 photo-like texture (view 1 as PNG bytes, then as a numpy array) and the same scene
169
+ moved by 24 px (view 2 as JPEG bytes). They have real image sizes, so the decoders, the resize, the readback
170
+ and the activations do the same work as for a photo.
171
+ - The warm-up does not keep results. Each later call computes its outputs from its own inputs.
172
+ The outputs are bit-identical with and without the warm-up (27 arrays of 3 reference pairs, plain, pose and
173
+ `predict_pairs`, compared between the old and the new code; `models/tests/test_warmup_device.py` also compares them).
174
+ - Without OpenCV, the `"pose"` warm-up captures the symmetric graph and skips the host PnP (an INFO log line tells you).
175
+
176
+ Cost (warm kernel cache, weights in the HF cache): the default variants add about 1.0 s of host warm-up
177
+ (`pair` 0.15 s, `pose` 0.7 s, `batch` 0.16 s) to the about 6 s device warm-up. `from_pretrained` takes about 9 s.
178
+ The trace region and the DRAM use do not change: the pair graph and the symmetric graph were already captured
179
+ at start-up before this API. `warmup_variants=["pair"]` does not capture the symmetric graph.
180
+ Then the first `return_pose=True` call captures it and imports OpenCV (about 0.6 s more on that call).
181
+
182
+ What the warm-up cannot prepare:
183
+
184
+ - Input image size has no limit, but it changes only the host resize (no device shape, no compile).
185
+ A new image size costs nothing extra on its first use.
186
+ - A variant that is not prepared pays its cost on its first call.
187
+ Measured without the host warm-up (`warmup_variants=()`, `test_warmup_device.py`): the first `pair` call took 74 ms
188
+ (later calls 44 ms). The first `return_pose=True` call took 870 ms (later calls 259 ms), because it captures the
189
+ symmetric graph and imports OpenCV.
190
+ - Code outside this package: DUSt3R's `global_aligner`, your own OpenCV or matplotlib calls.
191
+ - The first `from_pretrained` on a chip with an empty kernel cache compiles every kernel (about 67 s).
192
+
193
+ ### First-call measurements
194
+
195
+ Fresh process per row; each call gets a NEW real pair (crops of the demo photos at different windows and sizes,
196
+ read from PNG files), never the warm-up input. Median of 3 processes (8 for `pose` and `pairs`), Galaxy Blackhole chip,
197
+ ETH dispatch, 12x10 grid, warm kernel cache. The host is shared with other jobs, so a single call varies by about 10 ms.
198
+ "Same input" is the median of 3 later calls on the input of call 1. The PnP-RANSAC time depends on the data,
199
+ so that is the fair reference for call 1. Script: `code/tools_prof/first_call_bench.py`.
200
+
201
+ Before (HEAD `09c7d11`: device graphs warmed, host not warmed):
202
+
203
+ | variant | `from_pretrained` | 1st call | 2nd | 10th | steady median (calls 11-30) | same input | 1st-call penalty |
204
+ |---|---|---|---|---|---|---|---|
205
+ | `model(path, path)` | 8.0 s | 46.6 ms | 41.0 | 43.6 | 46.8 | 46.7 | -0.1 ms |
206
+ | `return_pose=True` | 7.9 s | 240.4 ms | 234.2 | 246.3 | 253.8 | 227.4 | +18.2 ms |
207
+ | `predict_pairs` (4 pairs) | 8.1 s | 126.7 ms | 123.0 | 124.1 | 124.8 | 129.2 | -2.5 ms |
208
+ | `load_images` + `inference` (3 views, 6 pairs) | 7.7 s | 219.1 ms | 227.4 | 220.5 | 225.9 | 225.9 | -6.6 ms |
209
+ | `preprocess="crop"` | 7.5 s | 43.5 ms | 40.9 | 48.7 | 49.9 | 40.4 | +3.1 ms |
210
+ | `model(ndarray, ndarray)` | 8.0 s | 53.3 ms | 42.8 | 52.6 | 51.6 | 55.5 | -1.4 ms |
211
+
212
+ After (default `warmup_variants`):
213
+
214
+ | variant | `from_pretrained` | 1st call | 2nd | 10th | steady median (calls 11-30) | same input | 1st call / steady median |
215
+ |---|---|---|---|---|---|---|---|
216
+ | `model(path, path)` | 9.2 s | 51.8 ms | 47.8 | 44.0 | 49.0 | 51.9 | 1.06 |
217
+ | `return_pose=True` | 9.0 s | 240.9 ms | 241.0 | 254.3 | 255.6 | 230.8 | 0.94 |
218
+ | `predict_pairs` (4 pairs) | 9.1 s | 132.4 ms | 126.8 | 122.2 | 124.2 | 128.2 | 1.07 |
219
+ | `load_images` + `inference` (3 views, 6 pairs) | 9.3 s | 229.0 ms | 236.9 | 219.8 | 222.3 | 216.9 | 1.03 |
220
+ | `preprocess="crop"` (not in the defaults) | 8.9 s | 55.7 ms | 46.2 | 55.0 | 50.7 | 53.8 | 1.10 |
221
+ | `model(ndarray, ndarray)` | 8.7 s | 65.3 ms | 53.0 | 64.6 | 59.4 | 70.3 | 1.10 |
222
+
223
+ The first call of every variant is within 10 % of the steady median. The steady state did not change:
224
+ the device call (`timing_ms['forward']`, 30 calls) is 27.0 ms before and 26.9 / 27.1 ms after, in back-to-back runs.
225
+ The pose first call has a timing breakdown that is the same as the later calls (focal vote about 15 ms, RANSAC 120-140 ms).
226
+ The `test_warmup_device.py` run shows the effect in one process: the first call is 0.92x (`pair`), 1.08x (`pose`)
227
+ and 0.99x (`batch`) of the same-input warm calls, against 1.70x, 3.36x and 1.10x without the host warm-up.
228
+
229
+ ## Console output
230
+
231
+ `from_pretrained(verbose=False)` (the default) keeps the console readable.
232
+ Before it imports ttnn, it sets `TT_LOGGER_LEVEL=error` (tt-metal's C++ log: no info lines such as
233
+ `act_block_h_override ...`, `Opening user mode device driver`, `JIT cache stats`) and `LOGURU_LEVEL=WARNING`
234
+ (ttnn's Python log: no `Initial ttnn.CONFIG` dump). Errors still print, and so do this package's warnings
235
+ (for example, the ETH-dispatch fallback). It changes a variable only if you did not set it,
236
+ and only if ttnn (or loguru) is not imported yet. The variables stay set in the process and its child processes.
237
+
238
+ - `verbose=True` (or `MAST3R_VERBOSE=1`) keeps tt-metal's normal log and prints this package's INFO log
239
+ (weights, warm-up times, ready line) on stderr.
240
+ - Set `TT_LOGGER_LEVEL` or `LOGURU_LEVEL` yourself for a different level.
241
+
242
  ## Lifetime
243
 
244
  Use `with Mast3rP150.from_pretrained() as model:` or call `model.close()`.
README.md CHANGED
@@ -38,6 +38,7 @@ In that environment, download this repo and install its `code/` directory:
38
  ```bash
39
  hf download changh95/mast3r-p150 --exclude "image/*" --local-dir mast3r-p150 && cd mast3r-p150
40
  pip install -e "code/[pose]" # the [pose] extra adds OpenCV for return_pose=True
 
41
  ```
42
 
43
  ```python
@@ -50,7 +51,12 @@ with Mast3rP150.from_pretrained(device_id=0) as model:
50
  out.save_ply("pointcloud.ply")
51
  ```
52
 
53
- The first call of `from_pretrained` downloads the NAVER weights to your HF cache. It also compiles the kernels and captures the metal trace (about 65 s). With a warm cache, startup takes about 7 s. The package does not install ttnn. You do not need the HTTP server.
 
 
 
 
 
54
 
55
  | | name | type / default | meaning |
56
  |---|---|---|---|
 
38
  ```bash
39
  hf download changh95/mast3r-p150 --exclude "image/*" --local-dir mast3r-p150 && cd mast3r-p150
40
  pip install -e "code/[pose]" # the [pose] extra adds OpenCV for return_pose=True
41
+ pip install -e "code/[pose,server,test]" # also the HTTP server and the tests (see PYTHON.md)
42
  ```
43
 
44
  ```python
 
51
  out.save_ply("pointcloud.ply")
52
  ```
53
 
54
+ The paths `media/source_1.png` and `media/source_2.png` are relative to the repo root.
55
+ Run the snippet from the repo root, or give absolute paths.
56
+
57
+ The first call of `from_pretrained` downloads the NAVER weights to your HF cache. It also compiles the kernels and captures the metal trace (about 65 s). The package does not install ttnn. You do not need the HTTP server.
58
+
59
+ `Mast3rP150.from_pretrained()` prepares the plain call, the pose call (`return_pose=True`) and `predict_pairs` before it returns. Thus the first call is as fast as the later calls (within 10 %). Start-up takes about 9 s with a warm kernel cache. Use `warmup_variants=["pair"]` for a faster start-up, or `model.warmup("crop")` to prepare more variants later.
60
 
61
  | | name | type / default | meaning |
62
  |---|---|---|---|
code/examples/quickstart.py CHANGED
@@ -4,7 +4,8 @@
4
 
5
  python examples/quickstart.py [--out-dir quickstart_out] [--pose]
6
 
7
- Reads media/source_1.png + media/source_2.png, runs the pair on the chip and writes
 
8
  <out-dir>/pointmaps.npz (the server's npz format), <out-dir>/pointcloud.ply (fused coloured
9
  cloud) and <out-dir>/depth.png (inputs + view-1 depth + confidence).
10
  """
 
4
 
5
  python examples/quickstart.py [--out-dir quickstart_out] [--pose]
6
 
7
+ Reads media/source_1.png + media/source_2.png (found relative to this file, so any working directory
8
+ works), runs the pair on the chip and writes (relative to the current directory)
9
  <out-dir>/pointmaps.npz (the server's npz format), <out-dir>/pointcloud.ply (fused coloured
10
  cloud) and <out-dir>/depth.png (inputs + view-1 depth + confidence).
11
  """
code/mast3r_p150/__init__.py CHANGED
@@ -11,8 +11,9 @@
11
  See PYTHON.md for the full API. ``import mast3r_p150`` needs no device and does not import ttnn.
12
  """
13
  from .inputs import load_images, make_pairs, to_pil
14
- from .model import DEFAULT_REPO_ID, DEFAULT_REVISION, Mast3rP150, PairResult, Pose
 
15
 
16
  __version__ = "0.1.0"
17
  __all__ = ["Mast3rP150", "PairResult", "Pose", "load_images", "make_pairs", "to_pil",
18
- "DEFAULT_REPO_ID", "DEFAULT_REVISION", "__version__"]
 
11
  See PYTHON.md for the full API. ``import mast3r_p150`` needs no device and does not import ttnn.
12
  """
13
  from .inputs import load_images, make_pairs, to_pil
14
+ from .model import (DEFAULT_REPO_ID, DEFAULT_REVISION, DEFAULT_WARMUP_VARIANTS, WARMUP_PRESETS, Mast3rP150,
15
+ PairResult, Pose)
16
 
17
  __version__ = "0.1.0"
18
  __all__ = ["Mast3rP150", "PairResult", "Pose", "load_images", "make_pairs", "to_pil",
19
+ "DEFAULT_REPO_ID", "DEFAULT_REVISION", "DEFAULT_WARMUP_VARIANTS", "WARMUP_PRESETS", "__version__"]
code/mast3r_p150/model.py CHANGED
@@ -16,6 +16,7 @@ from __future__ import annotations
16
  import io
17
  import logging
18
  import os
 
19
  import threading
20
  import time
21
  import warnings
@@ -36,7 +37,7 @@ from models.demos.mast3r.postprocess import (
36
  preprocess_mode,
37
  )
38
 
39
- from .inputs import ImageLike, View, prepare_view
40
 
41
  log = logging.getLogger("mast3r_p150")
42
 
@@ -73,6 +74,98 @@ def _activate(raw: torch.Tensor):
73
  return activate_pts3d(raw)[0].numpy(), activate_conf(raw)[0].numpy()
74
 
75
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
76
  # ---------------------------------------------------------------- outputs
77
 
78
  @dataclass
@@ -245,6 +338,7 @@ class _Engine:
245
  self.traced = bool(fused.enabled and fused.trace)
246
  self.lock = threading.RLock()
247
  self._pending = None
 
248
 
249
  def _inp(self, v: View) -> torch.Tensor:
250
  return v.pixels if self.uint8_input else v.float_input()
@@ -297,24 +391,52 @@ class _Engine:
297
  self._pending = None
298
  self.lock.release()
299
 
300
- def warmup(self, runs: int, capture_pose: bool) -> None:
301
- """The server's warm-up: ``runs`` forwards on a zero pair (the first compiles and captures),
302
- a capturing forward when needed, and the symmetric (pose) graph."""
303
  z = torch.zeros(1, 3, IMG_SIZE, IMG_SIZE, dtype=torch.uint8 if self.uint8_input else torch.float32)
304
  dummy = View(pixels=z, rgb=np.zeros((IMG_SIZE, IMG_SIZE, 3), np.uint8), record={})
305
  if not self.uint8_input:
306
  dummy._float = z
307
- for i in range(max(runs, 1 if self.traced else 0)):
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
308
  t0 = time.perf_counter()
309
- self.forward(dummy, dummy)
310
- log.info("warm-up forward %d: %.0f ms", i + 1, (time.perf_counter() - t0) * 1e3)
311
- if self.traced:
312
- if not self.port.get_model(self.state, self.device).trace_captured:
313
- raise RuntimeError("the warm-up did not capture a metal trace")
314
- if capture_pose and self.sym_graph:
315
  t0 = time.perf_counter()
316
- self.forward_pose(dummy, dummy)
317
- log.info("symmetric (pose) graph captured: %.0f ms", (time.perf_counter() - t0) * 1e3)
 
 
 
 
 
 
 
 
 
 
 
 
 
318
 
319
  def close(self) -> None:
320
  with self.lock:
@@ -355,6 +477,8 @@ class Mast3rP150:
355
  def __init__(self, engine, preprocess: str = "pad"):
356
  self._engine = engine
357
  self.preprocess = preprocess
 
 
358
  self._finalizer = weakref.finalize(self, _finalize, engine.close) if engine is not None else None
359
  Mast3rP150._open.add(self)
360
 
@@ -363,7 +487,8 @@ class Mast3rP150:
363
  def from_pretrained(cls, device_id: int = 0, *, device=None, dispatch: Optional[str] = None,
364
  repo_id: Optional[str] = None, revision: Optional[str] = None,
365
  weights_dir: Optional[str] = None, preprocess: Optional[str] = None,
366
- warmup: int = 1, capture_pose: bool = True) -> "Mast3rP150":
 
367
  """Load the weights, open the chip, build the device graph and capture its traces.
368
 
369
  ``device_id`` chip to open (``ttnn.open_device(device_id=...)``).
@@ -378,8 +503,20 @@ class Mast3rP150:
378
  (download on first use) at the pinned commit, or a local directory.
379
  Env: ``HF_MODEL``, ``TT_WEIGHTS_REVISION``, ``MAST3R_WEIGHTS_DIR``.
380
  ``preprocess`` default view preprocessing, ``"pad"`` or ``"crop"`` (``$MAST3R_PREPROC``).
381
- ``warmup`` warm-up forwards (the first compiles the kernels and captures the trace).
382
- ``capture_pose`` also capture the symmetric graph that ``return_pose=True`` uses.
 
 
 
 
 
 
 
 
 
 
 
 
383
 
384
  Device-graph knobs (``TT_FUSED``, ``MAST3R_OPT``, ``MAST3R_CQS``, ...) are read from the
385
  environment once, as in the server; the defaults are the published configuration.
@@ -392,6 +529,19 @@ class Mast3rP150:
392
  mode = preprocess or preprocess_mode()
393
  if mode not in ("pad", "crop"):
394
  raise ValueError(f"preprocess must be 'pad' or 'crop', got {mode!r}")
 
 
 
 
 
 
 
 
 
 
 
 
 
395
  torch.set_grad_enabled(False)
396
  t0 = time.perf_counter()
397
  ckpt = resolve_checkpoint_path(repo_id=repo_id or os.environ.get("HF_MODEL") or DEFAULT_REPO_ID,
@@ -414,14 +564,24 @@ class Mast3rP150:
414
  info.update(weights=str(ckpt), fused_opts=list(fused.opts), num_cqs=fused.num_cqs,
415
  traced=bool(fused.enabled and fused.trace))
416
  engine = _Engine(state, device, port, ttnn, owns, info)
 
417
  try:
418
- engine.warmup(warmup, capture_pose)
 
 
 
 
 
 
419
  except BaseException:
420
- engine.close()
 
 
 
421
  raise
422
  info["startup_s"] = round(time.perf_counter() - t0, 1)
423
  log.info("ready: dispatch %s, grid %s, %.1f s", info["dispatch"], info["grid"], info["startup_s"])
424
- return cls(engine, mode)
425
 
426
  @property
427
  def info(self) -> dict:
@@ -453,6 +613,80 @@ class Mast3rP150:
453
  raise RuntimeError("this Mast3rP150 is closed")
454
  return self._engine
455
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
456
  # ------------------------------------------------------------ inference
457
  def __call__(self, view1: ImageLike, view2: ImageLike, *, return_pose: bool = False,
458
  intrinsics: Optional[Sequence[Sequence[float]]] = None, conf_pct: float = 50.0,
 
16
  import io
17
  import logging
18
  import os
19
+ import sys
20
  import threading
21
  import time
22
  import warnings
 
37
  preprocess_mode,
38
  )
39
 
40
+ from .inputs import ImageLike, View, load_images, make_pairs, prepare_view
41
 
42
  log = logging.getLogger("mast3r_p150")
43
 
 
74
  return activate_pts3d(raw)[0].numpy(), activate_conf(raw)[0].numpy()
75
 
76
 
77
+ # ---------------------------------------------------------------- warm-up variants
78
+
79
+ #: Named warm-up variants. A variant is the set of call options that changes what runs on the first
80
+ #: call: ``return_pose`` (the symmetric device graph + OpenCV PnP), ``batch`` (the pipelined
81
+ #: ``predict_pairs`` / ``inference`` path) and ``preprocess`` (``pad`` / ``crop``; ``None`` = the model's).
82
+ WARMUP_PRESETS: Dict[str, Dict[str, Any]] = {
83
+ "pair": {"return_pose": False, "batch": False, "preprocess": None},
84
+ "pose": {"return_pose": True, "batch": False, "preprocess": None},
85
+ "batch": {"return_pose": False, "batch": True, "preprocess": None},
86
+ "crop": {"return_pose": False, "batch": False, "preprocess": "crop"},
87
+ }
88
+ #: ``from_pretrained(warmup_variants=...)`` default: the plain call, the pose call and the batch helpers.
89
+ DEFAULT_WARMUP_VARIANTS: Tuple[str, ...] = ("pair", "pose", "batch")
90
+ _VARIANT_KEYS = ("return_pose", "batch", "preprocess")
91
+
92
+
93
+ def normalize_warmup_variants(spec) -> List[Dict[str, Any]]:
94
+ """``warmup_variants`` -> list of variant dicts ``{return_pose, batch, preprocess}`` (duplicates removed).
95
+
96
+ ``spec``: ``None`` / ``()`` (no variant), ``"all"``, one name of :data:`WARMUP_PRESETS`, one dict of
97
+ call options, or a list of names / dicts (e.g. ``["pair", {"return_pose": True, "preprocess": "crop"}]``).
98
+ """
99
+ if spec is None:
100
+ return []
101
+ if isinstance(spec, (str, dict)):
102
+ spec = [spec]
103
+ out: List[Dict[str, Any]] = []
104
+ for item in spec:
105
+ if isinstance(item, str):
106
+ if item == "all":
107
+ out.extend(normalize_warmup_variants(list(WARMUP_PRESETS)))
108
+ continue
109
+ if item not in WARMUP_PRESETS:
110
+ raise ValueError(f"unknown warm-up variant {item!r}; use one of {sorted(WARMUP_PRESETS)}, "
111
+ "'all', or a dict of {return_pose, batch, preprocess}")
112
+ v = dict(WARMUP_PRESETS[item])
113
+ elif isinstance(item, dict):
114
+ bad = set(item) - set(_VARIANT_KEYS)
115
+ if bad:
116
+ raise ValueError(f"unknown warm-up variant option(s) {sorted(bad)}; allowed: {_VARIANT_KEYS}")
117
+ v = {"return_pose": bool(item.get("return_pose", False)), "batch": bool(item.get("batch", False)),
118
+ "preprocess": item.get("preprocess")}
119
+ else:
120
+ raise TypeError(f"a warm-up variant is a name or a dict, got {type(item).__name__}")
121
+ if v["preprocess"] not in (None, "pad", "crop"):
122
+ raise ValueError(f"warm-up preprocess must be 'pad', 'crop' or None, got {v['preprocess']!r}")
123
+ if v not in out:
124
+ out.append(v)
125
+ return out
126
+
127
+
128
+ def variant_name(v: Dict[str, Any]) -> str:
129
+ """Short label of a normalized variant (``pair``, ``pose``, ``batch``, ``pair/crop``, ...)."""
130
+ name = "pose" if v["return_pose"] else "batch" if v["batch"] else "pair"
131
+ if v["return_pose"] and v["batch"]:
132
+ name = "pose+batch"
133
+ return name + (f"/{v['preprocess']}" if v.get("preprocess") else "")
134
+
135
+
136
+ def warmup_images(seed: int = 0) -> Tuple[bytes, bytes, np.ndarray]:
137
+ """Realistic dummy inputs for the warm-up: a 640x480 photo-like texture (smooth random blobs plus
138
+ fine noise) as view 1 (PNG bytes) and the same scene shifted by 24 px as view 2 (JPEG bytes), plus
139
+ the view-1 pixels as a ``(480, 640, 3)`` uint8 array. Non-square, so ``pad`` and ``crop`` do real work;
140
+ the two decoders, the bicubic resize and the network see inputs of the same sizes as real photos."""
141
+ from PIL import Image
142
+ rng = np.random.default_rng(seed)
143
+ coarse = rng.integers(0, 256, size=(12, 17, 3), dtype=np.uint8)
144
+ base = np.asarray(Image.fromarray(coarse).resize((664, 480), Image.BICUBIC), dtype=np.int16)
145
+ base = np.clip(base + rng.integers(-12, 13, size=base.shape), 0, 255).astype(np.uint8)
146
+ v1, v2 = base[:, :640], base[:, 24:664]
147
+ b1, b2 = io.BytesIO(), io.BytesIO()
148
+ Image.fromarray(v1).save(b1, format="PNG")
149
+ Image.fromarray(v2).save(b2, format="JPEG", quality=92)
150
+ return b1.getvalue(), b2.getvalue(), np.ascontiguousarray(v1)
151
+
152
+
153
+ def quiet_native_logs(verbose: bool) -> Dict[str, str]:
154
+ """Before ttnn is imported: unless ``verbose``, make tt-metal's C++ logger (``TT_LOGGER_LEVEL``) print
155
+ errors only and ttnn's loguru logger (``LOGURU_LEVEL``) warnings and up. Values the caller already set win
156
+ (``os.environ.setdefault``); nothing changes once ttnn (or loguru) is imported. Returns what was set."""
157
+ if verbose or "ttnn" in sys.modules:
158
+ return {}
159
+ done = {}
160
+ for k, v in (("TT_LOGGER_LEVEL", "error"), ("LOGURU_LEVEL", "WARNING")):
161
+ if k == "LOGURU_LEVEL" and "loguru" in sys.modules:
162
+ continue
163
+ if k not in os.environ:
164
+ os.environ[k] = v
165
+ done[k] = v
166
+ return done
167
+
168
+
169
  # ---------------------------------------------------------------- outputs
170
 
171
  @dataclass
 
338
  self.traced = bool(fused.enabled and fused.trace)
339
  self.lock = threading.RLock()
340
  self._pending = None
341
+ self._eager_warm: set = set()
342
 
343
  def _inp(self, v: View) -> torch.Tensor:
344
  return v.pixels if self.uint8_input else v.float_input()
 
391
  self._pending = None
392
  self.lock.release()
393
 
394
+ def _dummy(self) -> View:
 
 
395
  z = torch.zeros(1, 3, IMG_SIZE, IMG_SIZE, dtype=torch.uint8 if self.uint8_input else torch.float32)
396
  dummy = View(pixels=z, rgb=np.zeros((IMG_SIZE, IMG_SIZE, 3), np.uint8), record={})
397
  if not self.uint8_input:
398
  dummy._float = z
399
+ return dummy
400
+
401
+ def graph_ready(self, pose: bool) -> bool:
402
+ """True when the device graph of this variant is compiled (and, when traced, captured)."""
403
+ if pose and not self.sym_graph:
404
+ pose = False # pose = two replays of the pair graph
405
+ if not self.traced:
406
+ return ("pose" if pose else "pair") in self._eager_warm
407
+ m = self.port.get_model(self.state, self.device)
408
+ return (m.input_kind(self._dummy().pixels) + ("_sym" if pose else "")) in m._kinds
409
+
410
+ def ensure_graph(self, pose: bool, runs: int = 1) -> float:
411
+ """Compile and capture the device graph of one variant on a zero pair (the server's warm-up input)
412
+ unless it is ready; idempotent. ``runs`` warm forwards for the pair graph (the first compiles and
413
+ captures). Returns the time spent in ms."""
414
+ if self.graph_ready(pose) and runs <= 1:
415
+ return 0.0
416
+ t_start = time.perf_counter()
417
+ dummy = self._dummy()
418
+ if pose and self.sym_graph:
419
  t0 = time.perf_counter()
420
+ self.forward_pose(dummy, dummy)
421
+ log.info("symmetric (pose) graph captured: %.0f ms", (time.perf_counter() - t0) * 1e3)
422
+ else:
423
+ for i in range(max(runs, 1)):
 
 
424
  t0 = time.perf_counter()
425
+ self.forward(dummy, dummy)
426
+ log.info("warm-up forward %d: %.0f ms", i + 1, (time.perf_counter() - t0) * 1e3)
427
+ if self.traced and not self.port.get_model(self.state, self.device).trace_captured:
428
+ raise RuntimeError("the warm-up did not capture a metal trace")
429
+ if not self.traced:
430
+ self._eager_warm.add("pose" if (pose and self.sym_graph) else "pair")
431
+ return (time.perf_counter() - t_start) * 1e3
432
+
433
+ def warmup(self, runs: int, capture_pose: bool) -> None:
434
+ """The server's warm-up: ``runs`` forwards on a zero pair (the first compiles and captures),
435
+ a capturing forward when needed, and the symmetric (pose) graph."""
436
+ if runs > 0 or self.traced:
437
+ self.ensure_graph(False, runs=runs)
438
+ if capture_pose and self.traced and self.sym_graph:
439
+ self.ensure_graph(True)
440
 
441
  def close(self) -> None:
442
  with self.lock:
 
477
  def __init__(self, engine, preprocess: str = "pad"):
478
  self._engine = engine
479
  self.preprocess = preprocess
480
+ self._warmed: List[Dict[str, Any]] = [] # normalized variants already warmed (host + device)
481
+ self._warm_lock = threading.Lock()
482
  self._finalizer = weakref.finalize(self, _finalize, engine.close) if engine is not None else None
483
  Mast3rP150._open.add(self)
484
 
 
487
  def from_pretrained(cls, device_id: int = 0, *, device=None, dispatch: Optional[str] = None,
488
  repo_id: Optional[str] = None, revision: Optional[str] = None,
489
  weights_dir: Optional[str] = None, preprocess: Optional[str] = None,
490
+ warmup: int = 1, capture_pose: Optional[bool] = None,
491
+ warmup_variants=DEFAULT_WARMUP_VARIANTS, verbose: Optional[bool] = None) -> "Mast3rP150":
492
  """Load the weights, open the chip, build the device graph and capture its traces.
493
 
494
  ``device_id`` chip to open (``ttnn.open_device(device_id=...)``).
 
503
  (download on first use) at the pinned commit, or a local directory.
504
  Env: ``HF_MODEL``, ``TT_WEIGHTS_REVISION``, ``MAST3R_WEIGHTS_DIR``.
505
  ``preprocess`` default view preprocessing, ``"pad"`` or ``"crop"`` (``$MAST3R_PREPROC``).
506
+ ``warmup`` warm-up forwards on a zero pair (the first compiles the kernels and captures the trace).
507
+ ``warmup_variants`` the call variants to prepare before this returns, so the FIRST real call of each
508
+ is as fast as the later ones: device graph compiled and captured, plus the host
509
+ first-use costs (image decoders, resize, thread pool, OpenCV import and PnP,
510
+ readback and activation buffers), run on realistic dummy images. Default
511
+ ``("pair", "pose", "batch")``; also ``"crop"``, ``"all"``, dicts of
512
+ ``{return_pose, batch, preprocess}``, or ``()`` for the device graph only.
513
+ :meth:`warmup` adds variants later. Outputs do not change.
514
+ ``capture_pose`` legacy switch: ``False`` drops ``"pose"`` from ``warmup_variants`` (no symmetric
515
+ graph is captured), ``True`` adds it.
516
+ ``verbose`` ``False`` (default; ``$MAST3R_VERBOSE=1`` sets ``True``): tt-metal's C++ log prints
517
+ errors only and ttnn's loguru log warnings and up (``TT_LOGGER_LEVEL`` /
518
+ ``LOGURU_LEVEL``, set only if unset and only before ttnn is first imported).
519
+ ``True``: tt-metal's normal log plus this package's INFO log on stderr.
520
 
521
  Device-graph knobs (``TT_FUSED``, ``MAST3R_OPT``, ``MAST3R_CQS``, ...) are read from the
522
  environment once, as in the server; the defaults are the published configuration.
 
529
  mode = preprocess or preprocess_mode()
530
  if mode not in ("pad", "crop"):
531
  raise ValueError(f"preprocess must be 'pad' or 'crop', got {mode!r}")
532
+ variants = normalize_warmup_variants(warmup_variants)
533
+ if capture_pose is False:
534
+ variants = [v for v in variants if not v["return_pose"]]
535
+ elif capture_pose is True and not any(v["return_pose"] for v in variants):
536
+ variants += normalize_warmup_variants("pose")
537
+ if verbose is None:
538
+ verbose = os.environ.get("MAST3R_VERBOSE", "0") not in ("", "0", "false", "False")
539
+ quiet_native_logs(verbose)
540
+ if verbose and not log.handlers:
541
+ h = logging.StreamHandler()
542
+ h.setFormatter(logging.Formatter("%(asctime)s mast3r_p150 %(levelname)s: %(message)s"))
543
+ log.addHandler(h)
544
+ log.setLevel(logging.INFO)
545
  torch.set_grad_enabled(False)
546
  t0 = time.perf_counter()
547
  ckpt = resolve_checkpoint_path(repo_id=repo_id or os.environ.get("HF_MODEL") or DEFAULT_REPO_ID,
 
564
  info.update(weights=str(ckpt), fused_opts=list(fused.opts), num_cqs=fused.num_cqs,
565
  traced=bool(fused.enabled and fused.trace))
566
  engine = _Engine(state, device, port, ttnn, owns, info)
567
+ model = None
568
  try:
569
+ t1 = time.perf_counter()
570
+ engine.warmup(warmup, any(v["return_pose"] for v in variants))
571
+ info["warmup_device_s"] = round(time.perf_counter() - t1, 2)
572
+ model = cls(engine, mode)
573
+ t1 = time.perf_counter()
574
+ model.warmup(variants)
575
+ info["warmup_host_s"] = round(time.perf_counter() - t1, 2)
576
  except BaseException:
577
+ if model is not None:
578
+ model.close()
579
+ else:
580
+ engine.close()
581
  raise
582
  info["startup_s"] = round(time.perf_counter() - t0, 1)
583
  log.info("ready: dispatch %s, grid %s, %.1f s", info["dispatch"], info["grid"], info["startup_s"])
584
+ return model
585
 
586
  @property
587
  def info(self) -> dict:
 
613
  raise RuntimeError("this Mast3rP150 is closed")
614
  return self._engine
615
 
616
+ # ------------------------------------------------------------ warm-up
617
+ @property
618
+ def warmed_variants(self) -> List[str]:
619
+ """Labels of the variants prepared so far (``pair``, ``pose``, ``batch``, ``pair/crop``, ...)."""
620
+ return [variant_name(v) for v in self._warmed]
621
+
622
+ def warmup(self, variants=None, *, runs: int = 2, **options) -> Dict[str, float]:
623
+ """Prepare call variants so that their first real call is fast; idempotent.
624
+
625
+ model.warmup() # the defaults: "pair", "pose", "batch"
626
+ model.warmup("crop") # a preset name (see WARMUP_PRESETS), or "all"
627
+ model.warmup(return_pose=True, preprocess="crop") # the options of one variant
628
+ model.warmup(["pair", {"batch": True}]) # several
629
+
630
+ For each variant not prepared yet: compile and capture its device graph (zero pair, as the
631
+ server does), then make ``runs`` complete calls through the public path on realistic dummy
632
+ images (:func:`warmup_images`: PNG + JPEG bytes, 640x480) so the host first-use costs are paid
633
+ here: PIL decoders, the pad/crop resize, the host thread pool, the readback and activation
634
+ buffers, and for ``return_pose`` the OpenCV import and one PnP-RANSAC. Nothing is cached:
635
+ a later call computes its outputs from its own inputs exactly as without the warm-up.
636
+ Returns ``{variant label: ms}`` for the variants prepared by this call (empty when all were ready).
637
+ """
638
+ if options:
639
+ if variants is not None:
640
+ raise TypeError("pass either variants or keyword options, not both")
641
+ variants = dict(options)
642
+ todo = normalize_warmup_variants(variants if variants is not None else DEFAULT_WARMUP_VARIANTS)
643
+ for v in todo: # explicit preprocess == the model's default -> same variant
644
+ if v["preprocess"] == self.preprocess:
645
+ v["preprocess"] = None
646
+ done: Dict[str, float] = {}
647
+ imgs = None
648
+ with self._warm_lock:
649
+ eng = self._eng()
650
+ for v in todo:
651
+ if v in self._warmed:
652
+ continue
653
+ t0 = time.perf_counter()
654
+ eng.ensure_graph(v["return_pose"])
655
+ if imgs is None:
656
+ imgs = warmup_images()
657
+ self._warm_host(v, max(int(runs), 0), imgs)
658
+ self._warmed.append(v)
659
+ done[variant_name(v)] = round((time.perf_counter() - t0) * 1e3, 1)
660
+ log.info("warm-up %s: %.0f ms", variant_name(v), done[variant_name(v)])
661
+ if done and self._engine is not None:
662
+ self._engine.info.setdefault("warmup_ms", {}).update(done)
663
+ return done
664
+
665
+ def _warm_host(self, v: Dict[str, Any], runs: int, imgs) -> None:
666
+ png, jpg, arr = imgs
667
+ mode = v["preprocess"] or self.preprocess
668
+ pose = v["return_pose"]
669
+ if pose:
670
+ try:
671
+ import cv2 # noqa: F401 (the PnP-RANSAC of the pose path imports it lazily)
672
+ except ImportError:
673
+ log.info("return_pose warm-up: OpenCV is not installed (pip install opencv-python-headless); "
674
+ "the symmetric device graph is ready, the host PnP is not warmed")
675
+ pose = False
676
+ for i in range(runs): # run 0: encoded PNG + JPEG; run 1: a numpy array + JPEG
677
+ a = png if i % 2 == 0 else arr
678
+ if v["batch"]:
679
+ if i == 0:
680
+ views = load_images([png, jpg], preprocess=mode)
681
+ self.inference(make_pairs(views), preprocess=mode) # 2 pairs, pipelined
682
+ else: # 4 pairs: the pipeline's steady state (prepare i+1 + finish i-1)
683
+ self.predict_pairs([(a, jpg), (jpg, a), (png, jpg), (jpg, png)], return_pose=pose,
684
+ preprocess=mode)
685
+ else:
686
+ r = self(a, jpg, return_pose=pose, preprocess=mode)
687
+ if i == 0:
688
+ r.point_cloud()
689
+
690
  # ------------------------------------------------------------ inference
691
  def __call__(self, view1: ImageLike, view2: ImageLike, *, return_pose: bool = False,
692
  intrinsics: Optional[Sequence[Sequence[float]]] = None, conf_pct: float = 50.0,
code/models/tests/test_api_host.py CHANGED
@@ -168,6 +168,10 @@ class FakeEngine:
168
  self._pending = None
169
  return self.forward(v1, v2)
170
 
 
 
 
 
171
  def close(self):
172
  self.closed = True
173
  self.device = None
@@ -308,3 +312,98 @@ def test_lifetime():
308
  with pytest.raises(RuntimeError, match="closed"):
309
  m(IMG1, IMG2)
310
  m.close() # idempotent
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
168
  self._pending = None
169
  return self.forward(v1, v2)
170
 
171
+ def ensure_graph(self, pose, runs=1):
172
+ self.calls.append(("graph", pose))
173
+ return 0.0
174
+
175
  def close(self):
176
  self.closed = True
177
  self.device = None
 
312
  with pytest.raises(RuntimeError, match="closed"):
313
  m(IMG1, IMG2)
314
  m.close() # idempotent
315
+
316
+
317
+ # ---------------------------------------------------------------- warm-up API
318
+
319
+ def test_warmup_variant_spec():
320
+ from mast3r_p150.model import DEFAULT_WARMUP_VARIANTS, WARMUP_PRESETS, normalize_warmup_variants, variant_name
321
+ assert DEFAULT_WARMUP_VARIANTS == ("pair", "pose", "batch") == mast3r_p150.DEFAULT_WARMUP_VARIANTS
322
+ assert [variant_name(v) for v in normalize_warmup_variants(DEFAULT_WARMUP_VARIANTS)] == ["pair", "pose", "batch"]
323
+ assert [variant_name(v) for v in normalize_warmup_variants("all")] == ["pair", "pose", "batch", "pair/crop"]
324
+ assert normalize_warmup_variants(None) == [] and normalize_warmup_variants(()) == []
325
+ assert normalize_warmup_variants("pose") == [WARMUP_PRESETS["pose"]]
326
+ assert normalize_warmup_variants({"return_pose": True, "preprocess": "crop"}) == \
327
+ [{"return_pose": True, "batch": False, "preprocess": "crop"}]
328
+ assert len(normalize_warmup_variants(["pair", {"return_pose": False}, "pair"])) == 1 # duplicates removed
329
+ with pytest.raises(ValueError, match="unknown warm-up variant"):
330
+ normalize_warmup_variants("huge")
331
+ with pytest.raises(ValueError, match="option"):
332
+ normalize_warmup_variants({"num_views": 3})
333
+ with pytest.raises(ValueError, match="preprocess"):
334
+ normalize_warmup_variants({"preprocess": "native"})
335
+ with pytest.raises(TypeError):
336
+ normalize_warmup_variants([3])
337
+
338
+
339
+ def test_warmup_images_are_realistic():
340
+ from mast3r_p150.model import warmup_images
341
+ png, jpg, arr = warmup_images()
342
+ a, b = Image.open(io.BytesIO(png)), Image.open(io.BytesIO(jpg))
343
+ assert a.format == "PNG" and b.format == "JPEG" and a.size == b.size == (640, 480)
344
+ assert arr.shape == (480, 640, 3) and arr.dtype == np.uint8 and arr.std() > 20 # textured, not constant
345
+ assert np.array_equal(np.asarray(a.convert("RGB")), arr)
346
+ png2, _, _ = warmup_images()
347
+ assert png2 == png # deterministic
348
+
349
+
350
+ def test_from_pretrained_signature():
351
+ import inspect
352
+ sig = inspect.signature(Mast3rP150.from_pretrained)
353
+ assert sig.parameters["warmup_variants"].default == ("pair", "pose", "batch")
354
+ assert sig.parameters["verbose"].default is None
355
+ assert sig.parameters["capture_pose"].default is None
356
+ assert "runs" in inspect.signature(Mast3rP150.warmup).parameters
357
+
358
+
359
+ def test_warmup_is_idempotent_and_does_not_change_outputs():
360
+ eng = FakeEngine()
361
+ m = Mast3rP150(eng)
362
+ before = m(IMG1, IMG2)
363
+ n0 = len(eng.calls)
364
+ done = m.warmup(["pair", "batch"])
365
+ assert list(done) == ["pair", "batch"] and all(t >= 0 for t in done.values())
366
+ assert m.warmed_variants == ["pair", "batch"]
367
+ assert ("graph", False) in eng.calls[n0:]
368
+ fwd = [c for c in eng.calls[n0:] if c[0] == "fwd"]
369
+ assert len(fwd) >= 4 # real calls through the public path
370
+ # the warm-up input is the textured dummy, not a cached real input
371
+ real_px = eng.calls[0][1]
372
+ assert not any(torch.equal(c[1], real_px) for c in fwd)
373
+ n1 = len(eng.calls)
374
+ assert m.warmup(["pair", "batch"]) == {} and m.warmup("pair") == {} and m.warmup(return_pose=False) == {}
375
+ assert len(eng.calls) == n1 # idempotent: no device work
376
+ assert m.warmup(preprocess="pad") == {} # pad == the model default == "pair"
377
+ after = m(IMG1, IMG2)
378
+ for k in ("pts3d1", "pts3d2", "conf1", "conf2", "rgb1", "rgb2"):
379
+ assert np.array_equal(getattr(before, k), getattr(after, k)), k
380
+ assert m.info["warmup_ms"].keys() == {"pair", "batch"}
381
+ assert list(m.warmup("crop")) == ["pair/crop"]
382
+ with pytest.raises(TypeError):
383
+ m.warmup("pair", return_pose=True)
384
+ m.close()
385
+ with pytest.raises(RuntimeError, match="closed"):
386
+ m.warmup("pose")
387
+
388
+
389
+ def test_warmup_pose_variant():
390
+ pytest.importorskip("cv2")
391
+ eng = FakeEngine()
392
+ m = Mast3rP150(eng)
393
+ assert list(m.warmup(return_pose=True)) == ["pose"]
394
+ assert ("graph", True) in eng.calls
395
+ import sys as _sys
396
+ assert "cv2" in _sys.modules
397
+ m.close()
398
+
399
+
400
+ def test_quiet_native_logs():
401
+ code = ("import os, sys; from mast3r_p150.model import quiet_native_logs as q; "
402
+ "os.environ.pop('TT_LOGGER_LEVEL', None); os.environ.pop('LOGURU_LEVEL', None); "
403
+ "assert q(True) == {} and 'TT_LOGGER_LEVEL' not in os.environ; "
404
+ "assert q(False) == {'TT_LOGGER_LEVEL': 'error', 'LOGURU_LEVEL': 'WARNING'}, os.environ; "
405
+ "os.environ['TT_LOGGER_LEVEL'] = 'debug'; os.environ.pop('LOGURU_LEVEL'); "
406
+ "assert q(False) == {'LOGURU_LEVEL': 'WARNING'} and os.environ['TT_LOGGER_LEVEL'] == 'debug'; "
407
+ "sys.modules['ttnn'] = object(); os.environ.pop('LOGURU_LEVEL'); assert q(False) == {}")
408
+ env = dict(os.environ, TT_VISIBLE_DEVICES="none", PYTHONPATH=_CODE_ROOT)
409
+ subprocess.run([sys.executable, "-c", code], check=True, env=env, cwd="/")
code/models/tests/test_warmup_device.py ADDED
@@ -0,0 +1,137 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-License-Identifier: Apache-2.0
2
+ """Device test of the warm-up API: the FIRST real call after ``from_pretrained`` is already fast
3
+ (needs a Blackhole chip, ttnn and the weights).
4
+
5
+ cd code && timeout -s INT 1200 python -m pytest -q -s models/tests/test_warmup_device.py
6
+
7
+ 1. A model opened with ``warmup_variants=()`` (device graph of the plain pair only) runs reference inputs:
8
+ the outputs, and the cold first-call cost of each variant (reported, not gated).
9
+ 2. A fresh model with the default ``warmup_variants`` (``pair``, ``pose``, ``batch``): the first call of each
10
+ variant on a NEW real input (crops of the demo photos, never the warm-up input) must be within
11
+ ``MAST3R_FIRST_TOL`` (10 %) + ``MAST3R_FIRST_SLACK_MS`` (5 ms, shared-host noise) of the median of 3 later
12
+ calls on the SAME input (the PnP-RANSAC time depends on the data, so the same input is the fair reference).
13
+ 3. The outputs of 2 are bit-identical to those of 1 (warm-up changes nothing), and ``warmup()`` is idempotent.
14
+ ``MAST3R_WARMUP_REPORT=<path>`` writes the numbers as JSON.
15
+ """
16
+ from __future__ import annotations
17
+
18
+ import json
19
+ import os
20
+ import statistics
21
+ import sys
22
+ import time
23
+
24
+ import numpy as np
25
+ import pytest
26
+
27
+ _CODE_ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", ".."))
28
+ if _CODE_ROOT not in sys.path:
29
+ sys.path.insert(0, _CODE_ROOT)
30
+
31
+ ttnn = pytest.importorskip("ttnn")
32
+ pytest.importorskip("cv2")
33
+
34
+ from PIL import Image # noqa: E402
35
+
36
+ from mast3r_p150 import Mast3rP150 # noqa: E402
37
+
38
+ _REPO = os.path.dirname(_CODE_ROOT)
39
+ TOL = float(os.environ.get("MAST3R_FIRST_TOL", "0.10"))
40
+ SLACK_MS = float(os.environ.get("MAST3R_FIRST_SLACK_MS", "5"))
41
+ REPEATS = 3
42
+
43
+
44
+ def _crops(tmp_path, n=6):
45
+ """n real pairs: the same window cut from both demo photos, written as PNG files."""
46
+ a0 = Image.open(os.path.join(_REPO, "media", "source_1.png")).convert("RGB")
47
+ b0 = Image.open(os.path.join(_REPO, "media", "source_2.png")).convert("RGB")
48
+ W, H = a0.size
49
+ out = []
50
+ for i in range(n):
51
+ f = 0.6 + 0.05 * i
52
+ w, h = int(W * f), int(H * (0.95 - 0.05 * i))
53
+ x, y = (W - w) * (i % 3) // 2, (H - h) * ((i + 1) % 3) // 2
54
+ pa, pb = tmp_path / f"r{i}_1.png", tmp_path / f"r{i}_2.png"
55
+ a0.crop((x, y, x + w, y + h)).save(pa)
56
+ b0.crop((x, y, x + w, y + h)).save(pb)
57
+ out.append((str(pa), str(pb)))
58
+ return out
59
+
60
+
61
+ def _variants(pairs):
62
+ """name -> callable(model) running one real call of that variant (each on its own new input)."""
63
+ return {
64
+ "pair": lambda m: m(*pairs[0]),
65
+ "pose": lambda m: m(*pairs[1], return_pose=True),
66
+ "batch": lambda m: m.predict_pairs([pairs[2], pairs[3]]),
67
+ }
68
+
69
+
70
+ def _arrays(res):
71
+ rs = res if isinstance(res, list) else [res]
72
+ out = []
73
+ for r in rs:
74
+ out += [r.raw1.numpy(), r.raw2.numpy(), r.pts3d1, r.conf2]
75
+ if r.pose is not None:
76
+ out += [r.pose.R, r.pose.t]
77
+ return out
78
+
79
+
80
+ def _timed(fn, m):
81
+ t0 = time.perf_counter()
82
+ r = fn(m)
83
+ return (time.perf_counter() - t0) * 1e3, r
84
+
85
+
86
+ def test_first_call_is_fast(tmp_path):
87
+ pairs = _crops(tmp_path)
88
+ calls = _variants(pairs)
89
+ report = {"tol": TOL, "slack_ms": SLACK_MS}
90
+
91
+ # 1. reference outputs + the cold first-call cost without the host/pose warm-up
92
+ t0 = time.perf_counter()
93
+ m = Mast3rP150.from_pretrained(device_id=0, warmup_variants=())
94
+ report["startup_no_variants_s"] = round(time.perf_counter() - t0, 2)
95
+ ref, cold = {}, {}
96
+ try:
97
+ for name, fn in calls.items():
98
+ first, r = _timed(fn, m)
99
+ ref[name] = _arrays(r)
100
+ same = statistics.median(_timed(fn, m)[0] for _ in range(REPEATS))
101
+ cold[name] = {"first_ms": round(first, 1), "same_input_warm_ms": round(same, 1)}
102
+ finally:
103
+ m.close()
104
+ report["cold"] = cold
105
+
106
+ # 2. default warm-up: the first real call of each variant is fast
107
+ t0 = time.perf_counter()
108
+ m = Mast3rP150.from_pretrained(device_id=0)
109
+ report["startup_default_s"] = round(time.perf_counter() - t0, 2)
110
+ report["warmup_ms"] = m.info.get("warmup_ms")
111
+ warm, failures = {}, []
112
+ try:
113
+ assert sorted(m.warmed_variants) == ["batch", "pair", "pose"]
114
+ for name, fn in calls.items():
115
+ first, r = _timed(fn, m)
116
+ got = _arrays(r)
117
+ same = statistics.median(_timed(fn, m)[0] for _ in range(REPEATS))
118
+ warm[name] = {"first_ms": round(first, 1), "same_input_warm_ms": round(same, 1),
119
+ "ratio": round(first / same, 3)}
120
+ # 3. bit-identical to the model without the warm-up
121
+ assert len(got) == len(ref[name]), name
122
+ for a, b in zip(got, ref[name]):
123
+ assert np.array_equal(a, b), f"{name}: outputs differ from the un-warmed model"
124
+ if first > same * (1.0 + TOL) + SLACK_MS:
125
+ failures.append(f"{name}: first {first:.1f} ms vs warm {same:.1f} ms")
126
+ t0 = time.perf_counter()
127
+ assert m.warmup() == {} # idempotent: everything is ready
128
+ assert (time.perf_counter() - t0) < 0.05
129
+ finally:
130
+ m.close()
131
+ report["warm"] = warm
132
+ print("\n# warm-up report: " + json.dumps(report, indent=1))
133
+ path = os.environ.get("MAST3R_WARMUP_REPORT")
134
+ if path:
135
+ with open(path, "w") as f:
136
+ json.dump(report, f, indent=1)
137
+ assert not failures, "first real call slower than the warm calls: " + "; ".join(failures)
code/pyproject.toml CHANGED
@@ -22,7 +22,7 @@ dependencies = [
22
  pose = ["opencv-python-headless>=4.9,<4.12"] # return_pose=True (PnP-RANSAC)
23
  viz = ["matplotlib"] # examples/quickstart.py figure
24
  server = ["fastapi", "uvicorn", "pydantic>=2"] # models/server/app.py
25
- test = ["pytest"]
26
 
27
  [project.urls]
28
  "Model card" = "https://huggingface.co/changh95/mast3r-p150"
 
22
  pose = ["opencv-python-headless>=4.9,<4.12"] # return_pose=True (PnP-RANSAC)
23
  viz = ["matplotlib"] # examples/quickstart.py figure
24
  server = ["fastapi", "uvicorn", "pydantic>=2"] # models/server/app.py
25
+ test = ["pytest", "httpx"] # host tests (fastapi TestClient needs httpx); add [server,pose] too
26
 
27
  [project.urls]
28
  "Model card" = "https://huggingface.co/changh95/mast3r-p150"
code/tools_prof/first_call_bench.py ADDED
@@ -0,0 +1,152 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ # SPDX-License-Identifier: Apache-2.0
3
+ """First-call latency of the Python API, one variant per FRESH process (needs a chip).
4
+
5
+ python tools_prof/first_call_bench.py --variant pair [--json out.json] [--dump out.npz] [--from-kw '{...}']
6
+
7
+ Variants (the request shapes a user hits): ``pair`` = model(path, path); ``pose`` = model(..., return_pose=True);
8
+ ``pairs`` = model.predict_pairs(4 pairs); ``inference`` = load_images(3 views) + make_pairs + model.inference;
9
+ ``crop`` = model(..., preprocess="crop"); ``array`` = model(np.ndarray, np.ndarray).
10
+
11
+ Times ``import mast3r_p150``, ``from_pretrained(**from_kw)``, then calls 1..10 and ``--steady`` more calls. EVERY call
12
+ gets a NEW real input pair (crops of the two demo photos at different windows/sizes, written once as PNG files under
13
+ ``--img-dir``), never the warm-up input. At the end the inputs of call 1 run 3 more times (warm): the
14
+ first-call penalty is call 1 minus the median of those (same input, so the data-dependent PnP-RANSAC time cancels). ``--dump`` stores the outputs of fixed reference pairs (demo pair + 2 crops,
15
+ plain and with pose) for a bit-identity check between code versions.
16
+ """
17
+ import argparse
18
+ import json
19
+ import os
20
+ import statistics
21
+ import sys
22
+ import time
23
+ from pathlib import Path
24
+
25
+ T_PROC = time.perf_counter()
26
+ CODE = Path(__file__).resolve().parents[1]
27
+ REPO = CODE.parent
28
+ sys.path.insert(0, str(CODE))
29
+
30
+ ap = argparse.ArgumentParser()
31
+ ap.add_argument("--variant", required=True, choices=["pair", "pose", "pairs", "inference", "crop", "array"])
32
+ ap.add_argument("--calls", type=int, default=10)
33
+ ap.add_argument("--steady", type=int, default=20)
34
+ ap.add_argument("--img-dir", default=os.environ.get("MAST3R_BENCH_IMGS", "/tmp/mast3r_first_call_imgs"))
35
+ ap.add_argument("--from-kw", default="{}", help="extra from_pretrained kwargs (JSON)")
36
+ ap.add_argument("--json", default=None)
37
+ ap.add_argument("--dump", default=None)
38
+ args = ap.parse_args()
39
+
40
+
41
+ def make_images(d: Path, n: int):
42
+ """n crop pairs of the demo photos (same window in both views), saved as PNG; deterministic."""
43
+ import numpy as np
44
+ from PIL import Image
45
+ d.mkdir(parents=True, exist_ok=True)
46
+ out = []
47
+ a0, b0 = Image.open(REPO / "media" / "source_1.png").convert("RGB"), Image.open(REPO / "media" / "source_2.png").convert("RGB")
48
+ rng = np.random.default_rng(1234)
49
+ for i in range(n):
50
+ pa, pb = d / f"c{i:03d}_1.png", d / f"c{i:03d}_2.png"
51
+ if not (pa.exists() and pb.exists()):
52
+ W, H = a0.size
53
+ f = rng.uniform(0.55, 0.95)
54
+ w, h = int(W * f), int(H * rng.uniform(0.55, 0.95))
55
+ x, y = int(rng.uniform(0, W - w)), int(rng.uniform(0, H - h))
56
+ flip = bool(rng.integers(2))
57
+ for im, p in ((a0, pa), (b0, pb)):
58
+ c = im.crop((x, y, x + w, y + h))
59
+ if flip:
60
+ c = c.transpose(Image.FLIP_LEFT_RIGHT)
61
+ c.save(p)
62
+ out.append((str(pa), str(pb)))
63
+ return out
64
+
65
+
66
+ need = (args.calls + args.steady) * (4 if args.variant == "pairs" else 2 if args.variant == "inference" else 1) + 4
67
+ IMGS = make_images(Path(args.img_dir), need)
68
+ REF = [(str(REPO / "media" / "source_1.png"), str(REPO / "media" / "source_2.png")), IMGS[-1], IMGS[-2]]
69
+ POOL = IMGS[:-4]
70
+
71
+ t0 = time.perf_counter()
72
+ import mast3r_p150 # noqa: E402
73
+ from mast3r_p150 import Mast3rP150 # noqa: E402
74
+ t_import = time.perf_counter() - t0
75
+ from_kw = json.loads(args.from_kw)
76
+ t0 = time.perf_counter()
77
+ model = Mast3rP150.from_pretrained(device_id=0, **from_kw)
78
+ t_from = time.perf_counter() - t0
79
+
80
+ it = iter(POOL)
81
+ _first = [] # the inputs of call 1, re-run warm at the end (same-input reference for the first call)
82
+
83
+
84
+ def one_call():
85
+ v = args.variant
86
+ if not _first:
87
+ need1 = {"pairs": 4, "inference": 2}.get(v, 1)
88
+ _first.extend(POOL[:need1])
89
+ if v == "pair":
90
+ a, b = next(it)
91
+ return model(a, b)
92
+ if v == "crop":
93
+ a, b = next(it)
94
+ return model(a, b, preprocess="crop")
95
+ if v == "array":
96
+ import numpy as np
97
+ from PIL import Image
98
+ a, b = next(it)
99
+ return model(np.asarray(Image.open(a).convert("RGB")), np.asarray(Image.open(b).convert("RGB")))
100
+ if v == "pose":
101
+ a, b = next(it)
102
+ return model(a, b, return_pose=True)
103
+ if v == "pairs":
104
+ return model.predict_pairs([next(it) for _ in range(4)])
105
+ if v == "inference":
106
+ p1, p2 = next(it), next(it)
107
+ views = mast3r_p150.load_images([p1[0], p1[1], p2[0]])
108
+ return model.inference(mast3r_p150.make_pairs(views))
109
+
110
+
111
+ lat, fwd = [], []
112
+ for i in range(args.calls + args.steady):
113
+ t0 = time.perf_counter()
114
+ r = one_call()
115
+ lat.append((time.perf_counter() - t0) * 1e3)
116
+ if hasattr(r, "timing_ms"):
117
+ fwd.append(r.timing_ms.get("forward"))
118
+ it = iter(_first * 3)
119
+ same = []
120
+ for _ in range(3):
121
+ t0 = time.perf_counter()
122
+ one_call()
123
+ same.append((time.perf_counter() - t0) * 1e3)
124
+ steady = lat[args.calls:]
125
+ med = statistics.median(steady) if steady else float("nan")
126
+ res = {"variant": args.variant, "from_kw": from_kw, "import_s": round(t_import, 2), "from_pretrained_s": round(t_from, 2),
127
+ "first_ms": round(lat[0], 1), "second_ms": round(lat[1], 1), "tenth_ms": round(lat[min(9, len(lat) - 1)], 1),
128
+ "steady_median_ms": round(med, 1), "steady_min_ms": round(min(steady), 1) if steady else None,
129
+ "first_over_median": round(lat[0] / med, 3),
130
+ "same_input_warm_ms": round(statistics.median(same), 1),
131
+ "first_penalty_ms": round(lat[0] - statistics.median(same), 1), "first_forward_ms": fwd[0] if fwd else None,
132
+ "lat_ms": [round(x, 1) for x in lat], "info": {k: v for k, v in model.info.items() if k != "open_kwargs"}}
133
+
134
+ if args.dump:
135
+ import numpy as np
136
+ arrs = {}
137
+ for k, (a, b) in enumerate(REF):
138
+ r = model(a, b)
139
+ arrs[f"r{k}_raw1"], arrs[f"r{k}_raw2"] = r.raw1.numpy(), r.raw2.numpy()
140
+ arrs[f"r{k}_pts1"], arrs[f"r{k}_conf2"] = r.pts3d1, r.conf2
141
+ rp = model(a, b, return_pose=True)
142
+ arrs[f"r{k}_pose_raw1"], arrs[f"r{k}_pose_pts2"] = rp.raw1.numpy(), rp.pts3d2
143
+ if rp.pose is not None:
144
+ arrs[f"r{k}_pose_R"], arrs[f"r{k}_pose_t"] = rp.pose.R, rp.pose.t
145
+ out = model.predict_pairs(REF)
146
+ for k, r in enumerate(out):
147
+ arrs[f"pp{k}_raw2"] = r.raw2.numpy()
148
+ np.savez(args.dump, **arrs)
149
+ model.close()
150
+ print("RESULT " + json.dumps({k: v for k, v in res.items() if k != "lat_ms"}))
151
+ if args.json:
152
+ Path(args.json).write_text(json.dumps(res, indent=1, default=str))