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)