superpoint-p150 / code /models /tests /test_api_host.py
changh95's picture
Python API (2026-10-04): pip install -e code/, from_pretrained() + model(...), Python-first quickstart
6ffd3f8 verified
Raw History Blame Contribute Delete
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
@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