Download code/models/tests/test_api_host.py from changh95/superpoint-p150: direct link, hf CLI and curl.
- Browser
- Download file 11.1 kB
-
https://huggingface.co/changh95/superpoint-p150/resolve/main/code/models/tests/test_api_host.py
- Command line
-
hf download hf://changh95/superpoint-p150/code/models/tests/test_api_host.py
-
curl -L -o test_api_host.py https://huggingface.co/changh95/superpoint-p150/resolve/main/code/models/tests/test_api_host.py
11.1 kB
| # 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 | |
| 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"}) | |
| 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) | |
| 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 | |