# SPDX-License-Identifier: Apache-2.0 """Host-only tests of the Python API (``tt_superpoint``): no device, no trace. cd code && TT_VISIBLE_DEVICES=none python -m pytest -q models/tests/test_api_host.py * package layout: ``tt_superpoint._port`` is the port package, its kernel directories exist; * input conversion (``to_plane``) gives the plane the HTTP server reads, for PIL / numpy / torch / path / bytes inputs, RGB / BGR / gray / RGBA / palette images; * parameter limits are the server's; * ``SuperPoint.__call__`` with a stand-in device model gives the same keypoints, scores and descriptors as the server's ``_predict_core`` driven by the same stand-in (order, scale); * list / batch / iterator calls keep the input order; ``close`` / ``with``. """ from __future__ import annotations import base64 import io import os from pathlib import Path import numpy as np import pytest import torch from PIL import Image import tt_superpoint from tt_superpoint import SuperPoint, SuperPointOutput, to_plane from tt_superpoint.inputs import is_single_image, split_batch from tt_superpoint.model import validate_params CODE = Path(__file__).resolve().parents[2] SAMPLE = CODE / "sample_data" / "house_in_field_1080p.jpg" def _server_plane(im: Image.Image) -> np.ndarray: """The server's decode path: ``_open_image`` (convert to RGB) + ``_r_plane``.""" from models.server import app as A rgb = im if im.mode == "RGB" else im.convert("RGB") return np.asarray(A._r_plane(rgb)) def test_package_layout(): from tt_superpoint._port.tt import fused_host, superpoint_ttnn # noqa: F401 - imports without a device from tt_superpoint._port.tt import conv_cell, nms_kernels assert fused_host.__name__ == "tt_superpoint._port.tt.fused_host" for kdir in (conv_cell._KDIR, nms_kernels._KDIR): assert os.path.isdir(kdir) and os.listdir(kdir), kdir assert tt_superpoint.__version__ # the port modules use relative imports only (they load under both package names) for py in (CODE / "models" / "tt").glob("*.py"): assert "from models." not in py.read_text(), py @pytest.mark.parametrize("mode", ["RGB", "L", "RGBA", "P", "LA", "I;16"]) def test_to_plane_pil_modes_match_server(mode): rng = np.random.default_rng(0) rgb = Image.fromarray(rng.integers(0, 256, (37, 53, 3), dtype=np.uint8), "RGB") if mode == "RGBA": im = rgb.convert("RGBA") im.putalpha(Image.fromarray(rng.integers(0, 256, (37, 53), dtype=np.uint8))) elif mode == "I;16": im = Image.fromarray(rng.integers(0, 65536, (37, 53), dtype=np.uint16)) else: im = rgb.convert(mode) np.testing.assert_array_equal(to_plane(im), _server_plane(im)) def test_to_plane_path_bytes_and_arrays_agree(): with Image.open(SAMPLE) as im: im.load() ref = _server_plane(im) rgb = np.asarray(im.convert("RGB")) assert ref.shape == (900, 1600) and ref.dtype == np.uint8 np.testing.assert_array_equal(to_plane(str(SAMPLE)), ref) np.testing.assert_array_equal(to_plane(SAMPLE), ref) np.testing.assert_array_equal(to_plane(SAMPLE.read_bytes()), ref) np.testing.assert_array_equal(to_plane(rgb), ref) # HWC RGB np.testing.assert_array_equal(to_plane(rgb[..., ::-1].copy(), bgr=True), ref) # HWC BGR (cv2) np.testing.assert_array_equal(to_plane(rgb[..., 0]), ref) # gray / plane np.testing.assert_array_equal(to_plane(torch.from_numpy(rgb.copy()).permute(2, 0, 1)), ref) # CHW uint8 f = torch.from_numpy(rgb).permute(2, 0, 1).float() / 255.0 # ToTensor() np.testing.assert_array_equal(to_plane(f), ref) np.testing.assert_array_equal(to_plane(f.numpy().transpose(1, 2, 0)), ref) # float HWC np.testing.assert_array_equal(to_plane(torch.from_numpy(rgb[..., 0]).long()), ref) # int64 plane def test_to_plane_errors(): with pytest.raises(ValueError): to_plane(np.full((8, 8), 1.5, dtype=np.float32)) with pytest.raises(ValueError): to_plane(np.full((8, 8), 300, dtype=np.int32)) with pytest.raises(ValueError): to_plane(np.zeros((8, 8, 7), dtype=np.uint8)) with pytest.raises(TypeError): to_plane(12345) def test_batch_detection(): a = np.zeros((4, 5, 3), np.uint8) assert is_single_image(a) and is_single_image("x.jpg") and is_single_image(torch.zeros(3, 4, 5)) assert not is_single_image([a, a]) and not is_single_image(np.zeros((2, 4, 5, 3), np.uint8)) assert len(split_batch(torch.zeros(3, 3, 4, 5))) == 3 def test_param_limits_match_server(): from models.server.app import PredictRequest ok = dict(max_keypoints=1024, keypoint_threshold=0.005, nms_radius=4, return_descriptors=True) assert validate_params(**ok) == ok for k, bad in (("max_keypoints", -2), ("max_keypoints", 480 * 640 + 1), ("keypoint_threshold", -0.1), ("keypoint_threshold", 1.5), ("nms_radius", -1), ("nms_radius", 33)): with pytest.raises(ValueError): validate_params(**{**ok, k: bad}) with pytest.raises(Exception): PredictRequest(image="", **{**ok, k: bad}) for k, edge in (("max_keypoints", -1), ("max_keypoints", 480 * 640), ("nms_radius", 0), ("nms_radius", 32), ("keypoint_threshold", 0.0), ("keypoint_threshold", 1.0)): validate_params(**{**ok, k: edge}) PredictRequest(image="", **{**ok, k: edge}) # ----------------------------------------------------------------------------- stand-in device model class _FakeTt: """Stands in for TtSuperPoint: deterministic keypoints from the plane contents (unsorted).""" device_resize = True kpc_ready = True nms_radius_traced = 4 border_removal_distance = 4 keypoint_threshold = 0.005 def __init__(self): self.released = 0 self.calls = [] def prepare_source(self, plane): return np.array(plane) def supports_device_nms_radius(self, r): return r is None or 1 <= int(r) <= 8 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, keypoint_threshold, max_keypoints, nms_radius, with_descriptors)) 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): self.released += 1 def _fake_model(): return SuperPoint(_FakeTt(), None, None, owns_device=False, dispatch=None, border_removal_distance=4, config={"model_id": "fake"}) @pytest.mark.parametrize("params", [{}, {"max_keypoints": 7}, {"nms_radius": 3, "keypoint_threshold": 0.01}, {"return_descriptors": False}, {"max_keypoints": -1}]) def test_call_matches_server_predict_core(params): from models.server import app as A model = _fake_model() out = model(str(SAMPLE), **params) assert isinstance(out, SuperPointOutput) req = A.PredictRequest(image=base64.b64encode(SAMPLE.read_bytes()).decode(), **params) A.STATE.update(ready=True, model=_FakeTt(), tt_in=None, fused=True, model_config={"border_removal_distance": 4}) try: resp = A.predict_dict(req) finally: A.STATE.clear() A.STATE["ready"] = False assert resp["num_keypoints"] == len(out) assert resp["keypoints"] == [[round(x, 3), round(y, 3)] for x, y in out.keypoints.double().tolist()] assert resp["scores"] == [round(s, 6) for s in out.scores.tolist()] assert out.image_size == (900, 1600) and out.scale == (2.5, 1.875) assert list(out.scores) == sorted(out.scores.tolist(), reverse=True) if params.get("return_descriptors", True): d = np.load(io.BytesIO(base64.b64decode(resp["descriptors"]["data"])))["descriptors"] np.testing.assert_array_equal(d, out.descriptors.half().numpy()) assert out.descriptors.dtype == torch.float32 and out.descriptors.shape == (len(out), 256) else: assert out.descriptors is None and "descriptors" not in resp assert out["keypoints"] is out.keypoints and set(out.numpy()) == {"keypoints", "scores", "descriptors"} assert out.to_dict()["num_keypoints"] == len(out) @pytest.mark.parametrize("num_workers", [0, 3]) def test_list_batch_and_iter_keep_order(num_workers): model = _fake_model() rng = np.random.default_rng(1) imgs = [rng.integers(0, 256, (90 + i, 160, 3), dtype=np.uint8) for i in range(7)] single = [model(im) for im in imgs] for got in (model(imgs, num_workers=num_workers), list(model.iter(iter(imgs), num_workers=num_workers))): assert len(got) == len(imgs) for a, b in zip(single, got): assert torch.equal(a.keypoints, b.keypoints) and torch.equal(a.descriptors, b.descriptors) assert a.image_size == b.image_size batch = np.stack([imgs[0][:90]] * 3) # (B, H, W, C) assert len(model(batch)) == 3 with pytest.raises(ValueError): model(imgs[0], nms_radius=40) with pytest.raises(TypeError): list(model.iter(imgs, foo=1)) def test_close_and_context_manager(): model = _fake_model() fake = model._m with model as m: assert m is model m(np.zeros((480, 640), np.uint8)) assert fake.released == 1 model.close() # idempotent assert fake.released == 1 with pytest.raises(RuntimeError): model(np.zeros((480, 640), np.uint8)) assert "closed" in repr(model) def test_dispatch_resolution(monkeypatch, tmp_path): from tt_superpoint import device as D patched = tmp_path / "patched" (patched / "tt_metal" / "impl" / "dispatch").mkdir(parents=True) (patched / "tt_metal" / "impl" / "dispatch" / "topology.cpp").write_text("single_chip_arch_1cq_no_dispatch_s") plain = tmp_path / "plain" (plain / "tt_metal").mkdir(parents=True) monkeypatch.delenv("SP_DISPATCH", raising=False) monkeypatch.setenv("TT_METAL_HOME", str(patched)) assert D.eth_dispatch_patch_present() and D.resolve_dispatch("auto") == "eth" monkeypatch.setenv("TT_METAL_HOME", str(plain)) with pytest.warns(RuntimeWarning, match="ETH-dispatch patch was not found"): assert D.resolve_dispatch("auto") == "worker" monkeypatch.setenv("SP_DISPATCH", "eth") assert D.resolve_dispatch("auto") == "eth" assert D.resolve_dispatch("worker") == "worker" with pytest.raises(ValueError): D.resolve_dispatch("tensix") def test_no_machine_specific_paths_in_runtime_code(): for sub in ("tt_superpoint", "models/tt", "models/server"): for py in (CODE / sub).rglob("*.py"): text = py.read_text() assert "/home/" not in text, py