File size: 2,493 Bytes
6ffd3f8
 
 
 
 
 
 
9350a1f
 
 
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
# SPDX-License-Identifier: Apache-2.0
"""Quickstart: the model-card Python snippet on the repo's demo image.

    pip install -e code/          # in an environment that has ttnn (tt-metal)
    python code/examples/quickstart.py [image] [--out-dir DIR] [--device-id N]

Writes <out-dir>/keypoints.png (keypoints drawn on the image), keypoints.json (keypoints +
scores) and descriptors.npy ((N, 256) float32). The demo image is found from the location of
this file, so the script runs from any directory; <out-dir> (default quickstart_out) is relative
to the current directory.
"""
import argparse
import json
from pathlib import Path

from PIL import Image, ImageDraw

from tt_superpoint import SuperPoint

DEMO = Path(__file__).resolve().parents[1] / "sample_data" / "house_in_field_1080p.jpg"


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("image", nargs="?", default=str(DEMO))
    ap.add_argument("--out-dir", default="quickstart_out")
    ap.add_argument("--device-id", type=int, default=0)
    ap.add_argument("--max-keypoints", type=int, default=1024)
    args = ap.parse_args()

    # --- model-card snippet ---------------------------------------------------------------
    with SuperPoint.from_pretrained(device_id=args.device_id) as model:
        out = model(args.image, max_keypoints=args.max_keypoints)
    print(len(out), "keypoints")
    print(out.keypoints[:3])      # (N, 2) [x, y] in original image pixels
    print(out.scores[:3])         # (N,) descending
    print(out.descriptors.shape)  # (N, 256) L2-normalised
    # ---------------------------------------------------------------------------------------

    out_dir = Path(args.out_dir)
    out_dir.mkdir(parents=True, exist_ok=True)
    im = Image.open(args.image).convert("RGB")
    draw = ImageDraw.Draw(im)
    r = max(2, round(max(im.size) / 400))
    for x, y in out.keypoints.tolist():
        draw.ellipse([x - r, y - r, x + r, y + r], outline=(255, 40, 40), width=max(1, r // 2))
    im.save(out_dir / "keypoints.png")
    (out_dir / "keypoints.json").write_text(json.dumps(
        {"image": str(args.image), "image_size": list(out.image_size), "num_keypoints": len(out),
         "keypoints": out.keypoints.tolist(), "scores": out.scores.tolist()}))
    import numpy as np

    np.save(out_dir / "descriptors.npy", out.descriptors.numpy())
    print("wrote", out_dir / "keypoints.png", out_dir / "keypoints.json", out_dir / "descriptors.npy")


if __name__ == "__main__":
    main()