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()
|