File size: 9,048 Bytes
6ffd3f8 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 | # SPDX-License-Identifier: Apache-2.0
"""Device tests of the Python API (``tt_superpoint.SuperPoint``). One chip, opened by the API.
cd code && python -m pytest -s -q models/tests/test_api_device.py
1. ``test_api_equals_server_and_timing``: ``model(...)`` returns the same keypoints, scores and
descriptors as the HTTP server's request path (``models.server.app.predict_dict`` /
``_predict_core``) on the demo JPEG, for 7 parameter sets and 6 input types; then the warm
``model(...)`` time against the server path in the same process and the published numbers.
2. ``test_api_on_external_device_equals_server_pipeline``: the server's own model build and
warm-up (``TtSuperPoint`` + ``_warmup_fused``) on a device opened by the caller, then
``SuperPoint.from_pretrained(device=...)`` on the same device: same output on the demo JPEG;
``close()`` leaves the caller's device open.
"""
from __future__ import annotations
import base64
import io
import os
import statistics
import time
from pathlib import Path
import numpy as np
import pytest
import torch
from PIL import Image
from tt_superpoint import SuperPoint
CODE = Path(__file__).resolve().parents[2]
SAMPLE = CODE / "sample_data" / "house_in_field_1080p.jpg"
DEVICE_ID = int(os.environ.get("TT_DEVICE_ID", "0"))
N_ITER = int(os.environ.get("SP_N_ITER", "200"))
PARAM_SETS = [
{},
{"max_keypoints": -1},
{"max_keypoints": 300, "keypoint_threshold": 0.015},
{"nms_radius": 3}, # precompiled per-radius device NMS variant
{"nms_radius": 0}, # host NMS fallback
{"nms_radius": 12, "max_keypoints": 2000}, # host NMS fallback, > 1024 keypoints
{"return_descriptors": False},
]
#: Published (OPT_REPORT.md round 10 / audit, ETH 12x10): served-like device_forward + postprocess
#: on the decoded 1600x900 JPEG (the rgb case adds the host R-channel copy of the HWC array), and
#: e2e_kpc_u8 on a 480x640 plane (H2D + trace + D2H + host decode; without the host sort by score).
PUBLISHED_MS = {"plane_1600x900": (1.58 + 0.15, 1.68 + 0.17), "rgb_1600x900": (1.58 + 0.15, 1.68 + 0.17),
"plane_480x640": (0.90, 0.96)}
def _same(out, resp, params):
assert resp["num_keypoints"] == len(out), (params, resp["num_keypoints"], len(out))
assert resp["keypoints"] == [[round(x, 3), round(y, 3)] for x, y in out.keypoints.double().tolist()], params
assert resp["scores"] == [round(s, 6) for s in out.scores.tolist()], params
if params.get("return_descriptors", True):
d = np.load(io.BytesIO(base64.b64decode(resp["descriptors"]["data"])))["descriptors"]
assert np.array_equal(d.view(np.uint16), out.descriptors.half().numpy().view(np.uint16)), params
else:
assert out.descriptors is None and "descriptors" not in resp
def _attach_server(A, model):
"""Point the server module at the API's device model (the server's own request path runs)."""
A.STATE.update(ready=True, model=model._m, tt_in=model._tt_in, fused=True, cfg={"fused": True},
model_config={"border_removal_distance": model.border_removal_distance})
def _detach_server(A):
A.STATE.clear()
A.STATE["ready"] = False
def _time(f, n):
ts = []
for _ in range(n):
t0 = time.perf_counter()
f()
ts.append((time.perf_counter() - t0) * 1e3)
return ts
def test_api_equals_server_and_timing():
from models.server import app as A
b64 = base64.b64encode(SAMPLE.read_bytes()).decode()
with Image.open(SAMPLE) as im:
im.load()
rgb_pil = im.convert("RGB")
rgb = np.asarray(rgb_pil)
plane = np.ascontiguousarray(rgb[..., 0])
with SuperPoint.from_pretrained(device_id=DEVICE_ID, precompile_nms_radii=(3,)) as model:
print(f"\nmodel: {model!r} config={model.config}")
_attach_server(A, model)
try:
# 1) equality with the server response, every parameter set
for params in PARAM_SETS:
resp = A.predict_dict(A.PredictRequest(image=b64, **params))
out = model(str(SAMPLE), **params)
_same(out, resp, params)
exp_dev_nms = params.get("nms_radius", 4) in range(1, 9)
assert out.device_nms == exp_dev_nms == resp["serving_path"]["device_nms"], params
print(f"same as server: params={params} n={len(out)} device_nms={out.device_nms}")
# 2) every input type -> the same output as the file path
ref = model(str(SAMPLE))
inputs = {
"bytes": SAMPLE.read_bytes(), "PIL": rgb_pil, "numpy HWC RGB": rgb,
"numpy HWC BGR": (rgb[..., ::-1].copy(), {"bgr": True}), "numpy plane": plane,
"torch CHW float": torch.from_numpy(rgb.copy()).permute(2, 0, 1).float() / 255.0,
}
for name, x in inputs.items():
x, kw = x if isinstance(x, tuple) else (x, {})
o = model(x, **kw)
assert torch.equal(o.keypoints, ref.keypoints) and torch.equal(o.scores, ref.scores), name
assert torch.equal(o.descriptors, ref.descriptors), name
outs = model([str(SAMPLE), rgb, plane])
assert all(torch.equal(o.descriptors, ref.descriptors) for o in outs)
# the 480x640 plane: the server's /predict_plane path
p480 = np.array(Image.fromarray(plane).resize((640, 480), Image.BILINEAR))
resp = A._predict_core(A.PlaneParams(), None, p480)
_same(model(p480), resp, {})
print(f"input types identical: {list(inputs)} + list call + 480x640 plane")
# 3) warm timing, model(...) vs the server's request path in the same process (alternating)
pp = A.PlaneParams()
cases = {
"plane_1600x900": (lambda: model(plane), lambda: A._predict_core(pp, None, plane)),
"rgb_1600x900": (lambda: model(rgb), lambda: A._predict_core(pp, None, plane)),
"plane_480x640": (lambda: model(p480), lambda: A._predict_core(pp, None, p480)),
"jpeg_path": (lambda: model(str(SAMPLE)),
lambda: A.predict_dict(A.PredictRequest(image=b64))),
}
rows = {}
for name, (fa, fs) in cases.items():
for f in (fa, fs):
_time(f, 20)
api, srv = [], []
for _ in range(4):
api += _time(fa, N_ITER // 4)
srv += _time(fs, N_ITER // 4)
rows[name] = (statistics.median(api), min(api), statistics.median(srv), min(srv))
pub = PUBLISHED_MS.get(name)
print(f"timing {name}: model() median {rows[name][0]:.3f} min {rows[name][1]:.3f} ms | "
f"server path median {rows[name][2]:.3f} min {rows[name][3]:.3f} ms"
+ (f" | published {pub[0]:.2f}-{pub[1]:.2f} ms" if pub else ""))
for name, (am, amin, sm, smin) in rows.items():
# same request path: not slower than the server path measured alongside (5 % + 50 us noise)
assert am <= sm * 1.05 + 0.05, (name, am, sm)
finally:
_detach_server(A)
def test_api_on_external_device_equals_server_pipeline():
import ttnn
from models.reference.superpoint_reference import load_reference_model
from models.server import app as A
from models.tt.superpoint_ttnn import TtSuperPoint
from tt_superpoint import DEFAULT_REVISION
from tt_superpoint.device import open_device
b64 = base64.b64encode(SAMPLE.read_bytes()).decode()
dev, mode = open_device(DEVICE_ID)
try:
# the server's own build + warm-up (models/server/app.py lifespan, fused path)
tm = load_reference_model(revision=DEFAULT_REVISION)
torch.set_grad_enabled(False)
sp = TtSuperPoint(tm, dev, input_height=A.INPUT_HEIGHT, input_width=A.INPUT_WIDTH, fused=True)
tt_in = sp.allocate_input(batch_size=1)
A.STATE.update(model=sp, tt_in=tt_in, fused=True, cfg={"fused": True},
model_config={"border_removal_distance": int(tm.config.border_removal_distance)})
A._warmup_fused(sp, tt_in, torch.zeros(1, 3, A.INPUT_HEIGHT, A.INPUT_WIDTH), ttnn, dev)
A.STATE["ready"] = True
resps = [A.predict_dict(A.PredictRequest(image=b64, **p)) for p in PARAM_SETS[:3]]
_detach_server(A)
ttnn.synchronize_device(dev)
sp.release()
ttnn.deallocate(tt_in)
del sp
with SuperPoint.from_pretrained(device=dev) as model:
for p, resp in zip(PARAM_SETS[:3], resps):
_same(model(str(SAMPLE), **p), resp, p)
print(f"\nexternal device ({mode}): API output == server pipeline output for {PARAM_SETS[:3]}")
g = dev.compute_with_storage_grid_size() # still open after close()
assert (g.x, g.y) == (12, 10) or mode == "worker"
finally:
_detach_server(A)
ttnn.close_device(dev)
|