superpoint-p150 / code /models /tests /test_api_warmup_host.py
changh95's picture
Python API: warm-up so the first real call is fast (warmup_variants, model.warmup()), quiet logs, install extras
9350a1f verified
Raw History Blame Contribute Delete
8.25 kB
# SPDX-License-Identifier: Apache-2.0
"""Host-only tests of the warm-up API (``warmup_variants``, ``model.warmup``): no device.
cd code && TT_VISIBLE_DEVICES=none python -m pytest -q models/tests/test_api_warmup_host.py
* the spec presets, dict overrides and the checks of the values;
* ``model.warmup(**variant)`` on a stand-in device model: the steps it runs, idempotency, the
thread pool that list calls reuse, no change of the outputs;
* the synthetic warm-up images and the decoder warm-up files;
* the log levels that ``verbose=False`` sets (and keeps when the user set them).
"""
from __future__ import annotations
import inspect
import sys
import numpy as np
import pytest
import torch
from tt_superpoint import SuperPoint, to_plane
from tt_superpoint import warmup as W
from tt_superpoint.model import _quiet_native_logs
def test_presets_and_overrides():
d = W.resolve_spec(None)
assert d == W.resolve_spec("default") == W.resolve_spec(True)
assert d["sizes"] == ((1920, 1080), (1600, 900), (1280, 720)) and d["nms_radii"] == (0, 4)
assert d["decoders"] == ("jpeg", "png") and d["num_workers"] == 4 and d["fallbacks"] is True
m = W.resolve_spec("minimal")
assert m == W.resolve_spec(False)
assert m["sizes"] == () and m["nms_radii"] == (4,) and not m["decoders"] and m["num_workers"] == 0
assert W.resolve_spec("all")["nms_radii"] == tuple(range(10))
o = W.resolve_spec({"nms_radii": [4, 3, 3], "sizes": [(1024, 768)], "decoders": ["JPG"]})
assert o["nms_radii"] == (3, 4) and o["sizes"] == ((1024, 768),) and o["decoders"] == ("jpeg",)
assert o["num_workers"] == 4 # other keys keep the defaults
@pytest.mark.parametrize("bad, exc", [
("fast", ValueError), ({"radius": 3}, ValueError), ({"nms_radii": [33]}, ValueError),
({"sizes": [(0, 10)]}, ValueError), ({"decoders": ["gif"]}, ValueError), ({"num_workers": -1}, ValueError),
(3, TypeError),
])
def test_bad_specs(bad, exc):
with pytest.raises(exc):
W.resolve_spec(bad)
def test_variant_kwargs_singular_and_plural():
s = W.variant_kwargs(size=(1024, 768), sizes=[(800, 600)], nms_radius=3, nms_radii=[5], decoder="png")
assert s["sizes"] == ((800, 600), (1024, 768)) and s["nms_radii"] == (3, 5) and s["decoders"] == ("png",)
assert s["num_workers"] == 0 and s["fallbacks"] is False
def test_api_signatures():
fp = inspect.signature(SuperPoint.from_pretrained).parameters
for name, default in (("warmup_variants", None), ("verbose", False), ("precompile_sizes", None),
("precompile_nms_radii", ())):
assert fp[name].default == default, name
wp = inspect.signature(SuperPoint.warmup).parameters
assert set(wp) == {"self", "size", "sizes", "nms_radius", "nms_radii", "decoder", "decoders", "num_workers",
"fallbacks"}
def test_synthetic_image_and_files():
a = W.synthetic_image(320, 200, seed=5, density=2.0)
assert a.shape == (200, 320, 3) and a.dtype == np.uint8
assert np.array_equal(a, W.synthetic_image(320, 200, seed=5, density=2.0)) # deterministic
assert not np.array_equal(a, W.synthetic_image(320, 200, seed=6, density=2.0))
assert a.std() > 20 # textured, not flat
for fmt in ("jpeg", "png", "bmp", "tiff"):
p = to_plane(W.encoded_image(fmt, 160, 120))
assert p.shape == (120, 160) and p.dtype == np.uint8, fmt
png = to_plane(W.encoded_image("png", 160, 120)) # lossless: the R plane of the synthetic image
assert np.array_equal(png, W.synthetic_image(160, 120, seed=7, density=2.0)[..., 0])
def test_keep_pillow_blocks(monkeypatch):
from PIL import Image
old = Image.core.get_blocks_max()
try:
monkeypatch.delenv("PILLOW_BLOCKS_MAX", raising=False)
Image.core.set_blocks_max(0)
assert W.keep_pillow_blocks(6) == 6 and Image.core.get_blocks_max() == 6
assert W.keep_pillow_blocks(2) == 6 # never lowered
Image.core.set_blocks_max(0)
monkeypatch.setenv("PILLOW_BLOCKS_MAX", "0") # the user's choice is kept
assert W.keep_pillow_blocks(6) == 0
finally:
Image.core.set_blocks_max(old)
def test_quiet_logs(monkeypatch):
monkeypatch.delitem(sys.modules, "ttnn", raising=False)
monkeypatch.delenv("TT_LOGGER_LEVEL", raising=False)
monkeypatch.setenv("LOGURU_LEVEL", "DEBUG")
_quiet_native_logs()
import os
assert os.environ["TT_LOGGER_LEVEL"] == "Error"
assert os.environ["LOGURU_LEVEL"] == "DEBUG" # set by the user: kept
# ----------------------------------------------------------------------------- stand-in device model
class _FakeTt:
"""Stands in for TtSuperPoint: deterministic outputs from the plane, records the warm-up steps."""
device_resize = True
kpc_ready = True
nms_radius_traced = 4
RSZ_MAX_VARIANTS = 8
def __init__(self):
self.calls, self.readback, self.resident = [], [], []
self._rsz_vars = {}
def prepare_source(self, plane):
h, w = plane.shape
if (h, w) != (480, 640):
self._rsz_vars[(w, h)] = object()
return np.array(plane)
def supports_device_resize(self, w, h):
return w <= 4096
def supports_device_nms_radius(self, r):
return r is None or 1 <= int(r) <= 8
def _variant(self, r):
from types import SimpleNamespace
return None if r in (None, 4) else SimpleNamespace(nms_map=f"nms_map_r{r}")
def warm_kpc_readback(self, r=None):
self.readback.append(r)
return 32
def _keypoints_from_resident(self, *a, **k):
self.resident.append(k)
def run_fused_keypoints_kpc(self, tt_in, host_in, *, keypoint_threshold, max_keypoints,
border_removal_distance, with_descriptors, nms_radius):
self.calls.append((host_in.shape, nms_radius, max_keypoints))
g = torch.Generator().manual_seed(int(host_in.sum()) % 1000 + int(nms_radius))
n = 50 if max_keypoints < 0 else min(50, max_keypoints)
kp = torch.randint(0, 480, (n, 2), generator=g).float()
sc = torch.rand(n, generator=g)
desc = torch.nn.functional.normalize(torch.randn(n, 256, generator=g), dim=1) if with_descriptors else None
return kp, sc, desc
def release(self):
pass
def _fake_model():
return SuperPoint(_FakeTt(), None, None, owns_device=False, dispatch=None, border_removal_distance=4,
config={"model_id": "fake"})
def test_warmup_steps_and_idempotency():
model = _fake_model()
fake = model._m
img = W.synthetic_image(800, 600, seed=1)
before = model(img)
spent = model.warmup(size=(1024, 768), nms_radius=3, decoder="jpeg", num_workers=2, fallbacks=True)
assert set(spent) == {"size/1024/768", "nms_radius/3", "decoder/jpeg", "num_workers/2", "fallbacks/3",
"fallbacks/4"}, spent
assert fake.readback == [3]
shapes = {c[0] for c in fake.calls}
assert (768, 1024) in shapes and (480, 640) in shapes
assert {c[1] for c in fake.calls} >= {3, 4}
assert any(c[2] == -1 for c in fake.calls) # fallbacks: max_keypoints=-1 on a dense image
assert len(fake.resident) == 2
n = len(fake.calls)
assert model.warmup(size=(1024, 768), nms_radius=3, decoder="jpeg", num_workers=2, fallbacks=True) == {}
assert len(fake.calls) == n # nothing ran again
# a radius warmed later also gets its fallbacks (fallbacks were requested before)
assert set(model.warmup(nms_radius=5)) == {"nms_radius/5", "fallbacks/5"}
after = model(img)
assert torch.equal(before.keypoints, after.keypoints) and torch.equal(before.descriptors, after.descriptors)
model.close()
def test_pool_kept_and_reused():
model = _fake_model()
model.warmup(num_workers=3)
pool = model._pool
assert pool is not None and len(pool._threads) == 3
imgs = [W.synthetic_image(64, 48, seed=i) for i in range(5)]
outs = model(imgs, num_workers=3)
assert model._pool is pool and len(outs) == 5
model(imgs, num_workers=6) # more workers: a larger pool replaces it
assert model._pool is not pool and model._pool._max_workers == 6
model.close()
assert model._pool is None