changh95's picture
Python API: warm-up so the first real call is fast (warmup_variants, model.warmup()), quiet logs, install extras
9350a1f verified
Raw History Blame Contribute Delete
33.1 kB
# SPDX-License-Identifier: Apache-2.0
"""``SuperPoint``: the Python entry point of the Tenstorrent Blackhole SuperPoint port.
from tt_superpoint import SuperPoint
with SuperPoint.from_pretrained(device_id=0) as model:
out = model("image.jpg")
out.keypoints # (N, 2) float32 [x, y] in original image pixels
out.scores # (N,) float32, descending
out.descriptors # (N, 256) float32, L2-normalised
The call runs the same request path as the HTTP server (``code/models/server/app.py``): one
metal trace (network, device NMS, keypoint list, bilinear descriptor sampling), the same
fallbacks, the same output order. Keypoints, scores and descriptors are equal to the server's
response for the same image and parameters (the server rounds the JSON floats and sends float16
descriptors).
"""
from __future__ import annotations
import logging
import threading
import time
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass, field
from typing import Any, Dict, Iterable, Iterator, List, Optional, Sequence, Tuple, Union
import numpy as np
import torch
from . import device as _device
from . import warmup as _warm
from .inputs import is_single_image, split_batch, to_plane
LOG = logging.getLogger("tt_superpoint")
DEFAULT_MODEL_ID = "magic-leap-community/superpoint"
#: The weights revision the published numbers and the server (serve.env) use.
DEFAULT_REVISION = "734450e9ffe229074f5998494ddc615475cdb20a"
#: Network frame (H, W). Every image is resized to it (bilinear, aspect ratio not kept).
INPUT_SIZE: Tuple[int, int] = (480, 640)
DESCRIPTOR_DIM = 256
#: Source sizes (W, H) whose device-resize variant is captured in from_pretrained (as serve.env).
DEFAULT_PRECOMPILE_SIZES: Tuple[Tuple[int, int], ...] = tuple(_warm.DEFAULT_WARMUP_VARIANTS["sizes"])
MAX_KEYPOINTS_CAP = INPUT_SIZE[0] * INPUT_SIZE[1]
MAX_NMS_RADIUS = 32
@dataclass
class SuperPointOutput:
"""Keypoints of one image. Same keys as the Hugging Face
``SuperPointImageProcessor.post_process_keypoint_detection`` result, so ``out["keypoints"]``
also works; keypoints stay float (sub-pixel scale) here instead of HF's int32.
Attributes:
keypoints: ``(N, 2)`` float32 tensor, ``[x, y]`` in pixels of the original image
(network-frame pixel position x ``scale``).
scores: ``(N,)`` float32 tensor, keypoint scores in (0, 1], sorted descending.
descriptors: ``(N, 256)`` float32 tensor, L2-normalised, in keypoint order; ``None``
when called with ``return_descriptors=False``.
image_size: ``(height, width)`` of the original image.
scale: ``(sx, sy)`` = original size / network size (640, 480).
device_nms: True when NMS ran on the device (False for the host NMS fallback, used for
``nms_radius`` 0 or > 8).
"""
keypoints: torch.Tensor
scores: torch.Tensor
descriptors: Optional[torch.Tensor]
image_size: Tuple[int, int]
scale: Tuple[float, float]
device_nms: bool = True
timing_ms: Dict[str, float] = field(default_factory=dict, repr=False)
def __len__(self) -> int:
return int(self.keypoints.shape[0])
def __getitem__(self, key: str):
if key not in ("keypoints", "scores", "descriptors"):
raise KeyError(key)
return getattr(self, key)
def keys(self):
return ("keypoints", "scores", "descriptors")
def numpy(self) -> Dict[str, Optional[np.ndarray]]:
"""``{"keypoints", "scores", "descriptors"}`` as numpy arrays (float32)."""
return {k: (None if getattr(self, k) is None else getattr(self, k).numpy()) for k in self.keys()}
def to_dict(self) -> Dict[str, Any]:
"""JSON-friendly dict (lists of floats), the same fields as the HTTP response."""
d = {
"num_keypoints": len(self),
"keypoints": self.keypoints.tolist(),
"scores": self.scores.tolist(),
"image_size": {"height": self.image_size[0], "width": self.image_size[1]},
"scale": {"x": self.scale[0], "y": self.scale[1]},
}
if self.descriptors is not None:
d["descriptors"] = self.descriptors.tolist()
return d
ImageLike = Any # path, bytes, PIL.Image, numpy array, torch tensor
class SuperPoint:
"""SuperPoint keypoint detector + descriptor on one Tenstorrent Blackhole chip.
Create it with :meth:`from_pretrained`; call it on an image (or a list of images); close it
with :meth:`close` or a ``with`` block. Device work is serialised by an internal lock, so one
instance can be shared between threads.
"""
def __init__(self, tt_model, tt_input, device, *, owns_device: bool, dispatch: Optional[str],
border_removal_distance: int, config: Dict[str, Any]):
self._m = tt_model
self._tt_in = tt_input
self._device = device
self._owns_device = owns_device
self._lock = threading.Lock()
self._closed = False
self.dispatch = dispatch
self.border_removal_distance = int(border_removal_distance)
#: Defaults of the per-call parameters (the server's defaults).
self.defaults = dict(max_keypoints=1024, keypoint_threshold=0.005, nms_radius=4, return_descriptors=True)
self.config = config
self._pool: Optional[ThreadPoolExecutor] = None # host decode threads of list calls (kept)
self._pool_lock = threading.Lock()
self._warm_done: set = set()
self._fallbacks_on = False
# ------------------------------------------------------------------ construction
@classmethod
def from_pretrained(
cls,
pretrained_model_name_or_path: str = DEFAULT_MODEL_ID,
*,
revision: Optional[str] = None,
device=None,
device_id: int = 0,
dispatch: str = "auto",
warmup_variants: Any = None,
precompile_sizes: Optional[Iterable[Tuple[int, int]]] = None,
precompile_nms_radii: Iterable[int] = (),
trace_region_size: int = _device.TRACE_REGION_SIZE,
local_files_only: bool = False,
verbose: bool = False,
) -> "SuperPoint":
"""Load the weights, open the device, build the model and capture its traces.
Args:
pretrained_model_name_or_path: Hugging Face repo id or a local directory with
``config.json`` + ``model.safetensors`` (default ``magic-leap-community/superpoint``).
revision: weights revision; default: the pinned revision of the published numbers
(``734450e9``) for the default repo, the repo's default branch otherwise.
device: an already-open ttnn device to use (it must have ``l1_small_size >= 32768``
and ``trace_region_size >= 32 MiB``; it is not closed by :meth:`close`). When
``None``, chip ``device_id`` is opened.
device_id: chip to open when ``device`` is None.
dispatch: ``"auto"`` (ETH dispatch + 12x10 grid when the patched tt-metal is present,
else the default dispatch with a warning), ``"eth"`` or ``"worker"``.
warmup_variants: the per-request variants to prepare now, so that the first real
call of each one is as fast as the later calls (see :mod:`tt_superpoint.warmup`):
``None`` / ``"default"`` (sizes 1920x1080, 1600x900, 1280x720 and 640x480,
``nms_radius`` 4, JPEG and PNG decoders, 4 list workers, the keypoint-list
fallbacks), ``"minimal"`` / ``False`` (640x480, radius 4 only), ``"all"`` (the
default plus every ``nms_radius`` 0..9), or a dict that replaces keys of the
default, e.g. ``{"nms_radii": (3, 4), "sizes": ((1024, 768),)}``.
precompile_sizes: replaces ``warmup_variants["sizes"]`` when given (older name).
precompile_nms_radii: added to ``warmup_variants["nms_radii"]`` (older name).
trace_region_size: trace region in bytes when this call opens the device.
local_files_only: load the weights from the HF cache only (no network).
verbose: False (default) sets the tt-metal and loguru log levels to errors and
warnings (``TT_LOGGER_LEVEL=Error``, ``LOGURU_LEVEL=WARNING``) when they are not
set and ttnn is not imported yet. True keeps the tt-metal info log.
Startup takes about 5 s with cached kernels (10-60 s when the kernels compile).
"""
if not verbose:
_quiet_native_logs()
spec = _warm.resolve_spec(warmup_variants)
if precompile_sizes is not None:
spec["sizes"] = _warm.normalize({"sizes": precompile_sizes})["sizes"]
if precompile_nms_radii:
spec["nms_radii"] = _warm.normalize({"nms_radii": tuple(spec["nms_radii"]) + tuple(precompile_nms_radii)})["nms_radii"]
torch.set_grad_enabled(False)
hf_model = _load_weights(pretrained_model_name_or_path, revision, local_files_only)
border = int(hf_model.config.border_removal_distance)
cfg = {"model_id": pretrained_model_name_or_path,
"revision": revision if revision is not None else
(DEFAULT_REVISION if pretrained_model_name_or_path == DEFAULT_MODEL_ID else None)}
import ttnn
from ._port.tt import fused_host as _fused
from ._port.tt.superpoint_ttnn import TtSuperPoint
owns = device is None
mode = None
if owns:
device, mode = _device.open_device(device_id, dispatch=dispatch, trace_region_size=trace_region_size)
tt_model = tt_in = self = None
try:
g = device.compute_with_storage_grid_size()
cfg["compute_grid"] = f"{g.x}x{g.y}"
cfg["dispatch"] = mode or "external"
tt_model = TtSuperPoint(hf_model, device, input_height=INPUT_SIZE[0], input_width=INPUT_SIZE[1],
fused=True, fused_stages=_fused.ALL_STAGES)
tt_in = tt_model.allocate_input(batch_size=1)
del hf_model
self = cls(tt_model, tt_in, device, owns_device=owns, dispatch=mode,
border_removal_distance=border, config=cfg)
self._warmup_base(ttnn)
t0 = time.perf_counter()
self.warmup(**spec)
self.config["warmup_s"]["variants"] = round(time.perf_counter() - t0, 2)
self.config["warmup_variants"] = spec
return self
except BaseException:
if self is not None:
self._shutdown_pool()
_release(device, tt_model, tt_in, close=owns)
raise
def _warmup_base(self, ttnn) -> None:
"""The server's warm-up (models/server/app.py ``_warmup_fused``): eager compile pass ->
trace capture -> one replay. Then the request path of the 480x640 network frame at the
default radius on synthetic images: readback buffers of every keypoint-list bucket, the
host post-processing (top-k, sort, descriptor normalisation) and the input conversions."""
m, tt_in, dev = self._m, self._tt_in, self._device
h, w = INPUT_SIZE
t = {}
with self._lock, torch.inference_mode():
t0 = time.perf_counter()
m.run_fused(tt_in, torch.zeros(1, 3, h, w, dtype=torch.float32)) # compiles every kernel
ttnn.synchronize_device(dev)
t["compile_forward"] = time.perf_counter() - t0
t0 = time.perf_counter()
m.capture_trace(tt_in, b=1)
t["trace_capture"] = time.perf_counter() - t0
res = m.run_fused(tt_in, torch.zeros(1, 3, h, w, dtype=torch.float32))
ttnn.synchronize_device(dev)
if not bool(torch.isfinite(res.descriptors_nchw).all()):
raise RuntimeError("warm-up forward produced non-finite outputs")
if m.trace_id is None:
raise RuntimeError("warm-up did not capture the metal trace")
t0 = time.perf_counter()
self._warm_radius(m.nms_radius_traced if m.kpc_ready else self.defaults["nms_radius"], (w, h))
# host first-use costs of the other call options and input types (no device work)
img = _warm.synthetic_image(64, 48, seed=3)
to_plane(torch.from_numpy(img).permute(2, 0, 1).float() / 255.0)
to_plane(img, bgr=True)
self._run_plane(to_plane(_warm.synthetic_image(w, h, seed=4, density=2.0)),
dict(self.defaults, max_keypoints=100, return_descriptors=False))
t["request_path"] = time.perf_counter() - t0
self.config["warmup_s"] = {k: round(v, 2) for k, v in t.items()}
def _warm_calls(self, size: Tuple[int, int], nms_radius: int, n: int = 3) -> None:
"""``n`` full calls on synthetic images of ``size`` (W, H) with ~650 keypoints each."""
p = dict(self.defaults, nms_radius=int(nms_radius))
for i in range(n):
self._run_plane(np.ascontiguousarray(_warm.synthetic_image(size[0], size[1], seed=11 + i, density=2.0)[..., 0]), p)
def _warm_radius(self, r: int, size: Tuple[int, int] = (INPUT_SIZE[1], INPUT_SIZE[0])) -> None:
m = self._m
if m.kpc_ready and m.supports_device_nms_radius(r):
with self._lock, torch.inference_mode():
m.warm_kpc_readback(r) # builds the radius variant when it is not the traced radius
self._warm_calls(size, r)
else:
self._warm_calls(size, r, n=2) # host NMS: torch only, nothing to build
def warmup(self, *, size: Optional[Tuple[int, int]] = None, sizes: Iterable[Tuple[int, int]] = (),
nms_radius: Optional[int] = None, nms_radii: Iterable[int] = (), decoder: Optional[str] = None,
decoders: Iterable[str] = (), num_workers: int = 0, fallbacks: bool = False) -> Dict[str, float]:
"""Prepare more per-request variants now (``from_pretrained`` already prepared the
``warmup_variants``). Idempotent: a variant that is ready is skipped.
Args:
size / sizes: source image sizes ``(width, height)``. A size in the device-resize range
gets its resize trace (approximately 10 ms each with cached kernels; at most 8 sizes
are kept, the least recently used one is released). Other sizes use the host resize.
nms_radius / nms_radii: 1..8 build the device NMS trace of that radius (approximately
50 ms each with cached kernels, 1-3 s when its kernels compile); 0 and > 8 warm the
host NMS.
decoder / decoders: image file formats to decode once: ``"jpeg"``, ``"png"``, ``"bmp"``,
``"webp"``, ``"tiff"``.
num_workers: start this many host threads for list / iterator calls.
fallbacks: warm the exact fallbacks of the keypoint list (more than 1024 NMS
candidates, candidate-slot overflow) for the prepared radii.
Returns:
seconds spent per step (an empty dict when everything was ready).
"""
self._check_open()
spec = _warm.variant_kwargs(size=size, sizes=sizes, nms_radius=nms_radius, nms_radii=nms_radii,
decoder=decoder, decoders=decoders, num_workers=num_workers, fallbacks=fallbacks)
m, h, w = self._m, INPUT_SIZE[0], INPUT_SIZE[1]
spent: Dict[str, float] = {}
def step(key, fn):
if key in self._warm_done:
return
t0 = time.perf_counter()
fn()
self._warm_done.add(key)
spent["/".join(str(k) for k in key)] = round(time.perf_counter() - t0, 3)
self._warm_done.add(("nms_radius", m.nms_radius_traced if m.kpc_ready else self.defaults["nms_radius"]))
for r in spec["nms_radii"]:
step(("nms_radius", r), lambda r=r: self._warm_radius(r))
if len(spec["sizes"]) > m.RSZ_MAX_VARIANTS:
LOG.warning("%d warm-up sizes, but only %d device-resize variants are kept", len(spec["sizes"]),
m.RSZ_MAX_VARIANTS)
for sw, sh in spec["sizes"]:
if (sh, sw) == INPUT_SIZE:
continue
on_device = m.device_resize and m.supports_device_resize(sw, sh)
if on_device and (sw, sh) in m._rsz_vars:
continue # trace resident (LRU), nothing to do
key = ("size", sw, sh)
if on_device:
self._warm_done.discard(key) # released by the LRU: build it again
step(key, lambda sw=sw, sh=sh: self._warm_calls((sw, sh), self.defaults["nms_radius"], n=2))
self._fallbacks_on |= spec["fallbacks"] # later device radii get their fallbacks warmed too
if self._fallbacks_on:
radii = sorted(r for k, r in (x for x in self._warm_done if x[0] == "nms_radius") if 1 <= r <= 8)
for r in radii:
step(("fallbacks", r), lambda r=r: self._warm_fallbacks(r))
if spec["decoders"]:
# freed Pillow image buffers are kept for the next decode (2 per decoding thread)
_warm.keep_pillow_blocks(2 * (1 + max(spec["num_workers"], self._pool_workers())))
if spec["num_workers"] > 0:
step(("num_workers", spec["num_workers"]), lambda: self._warm_pool(spec["num_workers"], spec["decoders"]))
# last: the large tensors of the steps above can shrink the heap the decodes grow
big = self._largest_warm_size()
for fmt in spec["decoders"]:
step(("decoder", fmt), lambda fmt=fmt: _decode_n(_warm.encoded_image(fmt, *big)))
if spent:
# end on the default request path: the steps above leave the host caches cold for it
self._warm_calls((INPUT_SIZE[1], INPUT_SIZE[0]), self.defaults["nms_radius"], n=2)
return spent
def _warm_fallbacks(self, r: int) -> None:
"""More than 1024 NMS candidates on a dense synthetic image (host top-k + sampling trace),
then the host post-processing of the resident maps (the slot-overflow fallback)."""
m = self._m
img = np.ascontiguousarray(_warm.synthetic_image(INPUT_SIZE[1], INPUT_SIZE[0], seed=21, density=6.0)[..., 0])
self._run_plane(img, dict(self.defaults, nms_radius=r))
self._run_plane(img, dict(self.defaults, nms_radius=r, max_keypoints=-1))
if m.kpc_ready:
var = m._variant(r)
with self._lock, torch.inference_mode():
m._keypoints_from_resident(self.defaults["keypoint_threshold"], self.defaults["max_keypoints"],
self.border_removal_distance, True,
None if var is None else var.nms_map)
def _pool_workers(self) -> int:
return 0 if self._pool is None else int(self._pool._max_workers)
def _get_pool(self, num_workers: int) -> ThreadPoolExecutor:
with self._pool_lock:
if self._pool is None or self._pool._max_workers < num_workers:
old, self._pool = self._pool, ThreadPoolExecutor(max_workers=num_workers,
thread_name_prefix="tt_superpoint_io")
if old is not None:
old.shutdown(wait=False)
return self._pool
def _largest_warm_size(self) -> Tuple[int, int]:
sizes = [(k[1], k[2]) for k in self._warm_done if k[0] == "size"] + [(INPUT_SIZE[1], INPUT_SIZE[0])]
return max(sizes, key=lambda s: s[0] * s[1])
def _warm_pool(self, num_workers: int, decoders: Sequence[str] = ()) -> None:
"""Start the host threads of list calls; each thread decodes one image file of the largest
warm size (its malloc arena grows once), then one list call on synthetic images."""
pool = self._get_pool(num_workers)
data = _warm.encoded_image(decoders[0] if decoders else "jpeg", *self._largest_warm_size())
barrier = threading.Barrier(num_workers + 1)
def task():
barrier.wait(30) # one task on each thread
_decode_n(data)
futs = [pool.submit(task) for _ in range(num_workers)]
barrier.wait(30)
for f in futs:
f.result()
imgs = [_warm.synthetic_image(INPUT_SIZE[1], INPUT_SIZE[0], seed=31 + i, density=2.0) for i in range(2)]
list(self.iter(imgs, num_workers=num_workers))
def _shutdown_pool(self) -> None:
with self._pool_lock:
if self._pool is not None:
self._pool.shutdown(wait=True)
self._pool = None
# ------------------------------------------------------------------ inference
def _check_open(self) -> None:
if self._closed:
raise RuntimeError("SuperPoint model is closed")
def _infer(self, plane: np.ndarray, *, max_keypoints: int, keypoint_threshold: float, nms_radius: int,
return_descriptors: bool):
"""The server's ``_infer_fused`` on an (H, W) uint8 plane at the source size ->
(kp (N, 2) network-frame xy, scores (N,), desc (N, 256) | None, device_nms)."""
from ._port.tt import postprocess as _post
m, tt_in, border = self._m, self._tt_in, self.border_removal_distance
if not plane.flags.writeable and plane.shape == INPUT_SIZE:
plane = np.array(plane) # torch.as_tensor warns on read-only arrays
host_in = m.prepare_source(plane)
if m.kpc_ready and m.supports_device_nms_radius(nms_radius):
with self._lock, torch.inference_mode():
kp, sc, desc = m.run_fused_keypoints_kpc(
tt_in, host_in, keypoint_threshold=keypoint_threshold, max_keypoints=max_keypoints,
border_removal_distance=border, with_descriptors=return_descriptors, nms_radius=nms_radius,
)
return kp, sc, desc, True
with self._lock, torch.inference_mode():
res = m.run_fused_prepared(tt_in, host_in, nms_radius=nms_radius)
with torch.inference_mode():
if res.nms_map is not None:
kp, sc, desc = _post.postprocess_from_nms_map(
res.nms_map, res.descriptors_nchw, keypoint_threshold=keypoint_threshold,
max_keypoints=max_keypoints, border_removal_distance=border, with_descriptors=return_descriptors,
)[0]
return kp, sc, desc, True
kp, sc, desc = _post.postprocess_keypoints(
res.scores_nchw, res.descriptors_nchw, nms_radius=nms_radius, keypoint_threshold=keypoint_threshold,
max_keypoints=max_keypoints, border_removal_distance=border, with_descriptors=return_descriptors,
)[0]
return kp, sc, desc, False
def _params(self, max_keypoints, keypoint_threshold, nms_radius, return_descriptors) -> Dict[str, Any]:
p = dict(self.defaults)
if max_keypoints is not None:
p["max_keypoints"] = int(max_keypoints)
if keypoint_threshold is not None:
p["keypoint_threshold"] = float(keypoint_threshold)
if nms_radius is not None:
p["nms_radius"] = int(nms_radius)
if return_descriptors is not None:
p["return_descriptors"] = bool(return_descriptors)
return validate_params(**p)
def _run_plane(self, plane: np.ndarray, p: Dict[str, Any], t_pre: float = 0.0) -> SuperPointOutput:
self._check_open()
orig_h, orig_w = plane.shape
t0 = time.perf_counter()
kp, sc, desc, device_nms = self._infer(plane, **p)
t1 = time.perf_counter()
order = torch.argsort(sc, descending=True) # same order as the server response
kp, sc = kp[order], sc[order]
if desc is not None:
desc = torch.index_select(desc, 0, order).float()
sx, sy = orig_w / INPUT_SIZE[1], orig_h / INPUT_SIZE[0]
kp_orig = (kp.double() * torch.tensor([sx, sy], dtype=torch.float64)).float()
t2 = time.perf_counter()
return SuperPointOutput(
keypoints=kp_orig, scores=sc.float(), descriptors=desc if p["return_descriptors"] else None,
image_size=(int(orig_h), int(orig_w)), scale=(sx, sy), device_nms=bool(device_nms),
timing_ms={"preprocess": round(t_pre * 1e3, 3), "device_forward": round((t1 - t0) * 1e3, 3),
"postprocess": round((t2 - t1) * 1e3, 3)},
)
def __call__(
self,
images: Union[ImageLike, Sequence[ImageLike]],
*,
max_keypoints: Optional[int] = None,
keypoint_threshold: Optional[float] = None,
nms_radius: Optional[int] = None,
return_descriptors: Optional[bool] = None,
bgr: bool = False,
num_workers: int = 4,
) -> Union[SuperPointOutput, List[SuperPointOutput]]:
"""Detect keypoints and compute descriptors.
Args:
images: one image, or a list / tuple of images, or a 4-D batch array / tensor
((B, H, W, C) numpy, (B, C, H, W) torch). An image is a file path, encoded
bytes, a ``PIL.Image``, a numpy array or a torch tensor of any size (see
:func:`tt_superpoint.inputs.to_plane`). The model reads channel 0 (R; the gray
plane of a grayscale image), resized to 480x640.
max_keypoints: keep the top-k by score (default 1024); -1 keeps every keypoint above
the threshold.
keypoint_threshold: minimum score after NMS, in [0, 1] (default 0.005).
nms_radius: NMS radius in network-frame pixels, 0..32 (default 4; 1..8 run on the
device, 0 and > 8 on the host).
return_descriptors: compute the 256-d descriptors (default True).
bgr: arrays are BGR (OpenCV); the R channel is then the last one.
num_workers: host threads that decode / convert the next images of a list while the
device runs the current one (lists only; 0 = no threads).
Returns:
one :class:`SuperPointOutput` for one image; a list of them, in input order, for a
list or a batch.
"""
p = self._params(max_keypoints, keypoint_threshold, nms_radius, return_descriptors)
self._check_open()
if is_single_image(images):
t0 = time.perf_counter()
plane = to_plane(images, bgr=bgr)
return self._run_plane(plane, p, time.perf_counter() - t0)
return list(self.iter(split_batch(images), bgr=bgr, num_workers=num_workers, **p))
def iter(self, images: Iterable[ImageLike], *, bgr: bool = False, num_workers: int = 4,
**params) -> Iterator[SuperPointOutput]:
"""Yield one :class:`SuperPointOutput` per image of ``images`` (any iterable, e.g. a
generator over a video or a directory), in order. Up to ``num_workers`` images are
decoded / converted on host threads ahead of the device; the device runs one image at a
time, so the outputs are the same as calling the model on each image. Keyword arguments:
the per-call parameters of :meth:`__call__`."""
p = self._params(params.pop("max_keypoints", None), params.pop("keypoint_threshold", None),
params.pop("nms_radius", None), params.pop("return_descriptors", None))
if params:
raise TypeError(f"unexpected keyword argument(s): {sorted(params)}")
def load(img):
t0 = time.perf_counter()
return to_plane(img, bgr=bgr), time.perf_counter() - t0
if num_workers <= 0:
for img in images:
plane, t = load(img)
yield self._run_plane(plane, p, t)
return
it = iter(images)
ex = self._get_pool(num_workers) # kept between calls (started by the warm-up)
pending = []
try:
for img in it:
pending.append(ex.submit(load, img))
if len(pending) > num_workers:
break
while pending:
plane, t = pending.pop(0).result()
nxt = next(it, _END)
if nxt is not _END:
pending.append(ex.submit(load, nxt))
yield self._run_plane(plane, p, t)
finally:
for f in pending:
f.cancel()
# ------------------------------------------------------------------ info / lifetime
@property
def device(self):
"""The ttnn device the model runs on."""
return self._device
@property
def input_size(self) -> Tuple[int, int]:
"""Network frame (height, width) = (480, 640)."""
return INPUT_SIZE
def close(self) -> None:
"""Release the traces and device tensors, and close the device if this object opened it.
Safe to call more than once."""
if self._closed:
return
self._closed = True
self._shutdown_pool()
with self._lock:
_release(self._device, self._m, self._tt_in, close=self._owns_device)
self._m = self._tt_in = None
def __enter__(self) -> "SuperPoint":
return self
def __exit__(self, *exc) -> None:
self.close()
def __repr__(self) -> str:
state = "closed" if self._closed else f"grid {self.config.get('compute_grid')}, dispatch {self.config.get('dispatch')}"
return f"SuperPoint({self.config.get('model_id')!r}, {state})"
_END = object()
def _decode_n(data: bytes, n: int = 3) -> None:
"""Decode an image file ``n`` times. The first decode imports the Pillow plug-in. The next
ones grow the malloc heap of this thread to the image size: glibc serves the first large
buffers with mmap and raises its mmap threshold only when such a buffer is freed."""
for _ in range(n):
to_plane(data)
def _quiet_native_logs() -> None:
"""Errors and warnings only from tt-metal (C++) and loguru (ttnn Python), unless the user set
the levels. The levels are read when ttnn is imported, so this has no effect after that."""
import os
import sys
if "ttnn" in sys.modules:
return
os.environ.setdefault("TT_LOGGER_LEVEL", "Error")
os.environ.setdefault("LOGURU_LEVEL", "WARNING")
def validate_params(*, max_keypoints: int, keypoint_threshold: float, nms_radius: int,
return_descriptors: bool) -> Dict[str, Any]:
"""Range checks of the per-call parameters (the HTTP server's limits)."""
if not -1 <= max_keypoints <= MAX_KEYPOINTS_CAP:
raise ValueError(f"max_keypoints must be in -1..{MAX_KEYPOINTS_CAP}, got {max_keypoints}")
if not 0.0 <= keypoint_threshold <= 1.0:
raise ValueError(f"keypoint_threshold must be in [0, 1], got {keypoint_threshold}")
if not 0 <= nms_radius <= MAX_NMS_RADIUS:
raise ValueError(f"nms_radius must be in 0..{MAX_NMS_RADIUS}, got {nms_radius}")
return dict(max_keypoints=max_keypoints, keypoint_threshold=keypoint_threshold, nms_radius=nms_radius,
return_descriptors=return_descriptors)
def _load_weights(name_or_path: str, revision: Optional[str], local_files_only: bool):
"""HF ``SuperPointForKeypointDetection`` (fp32 reference whose weights the device model
uses), from the HF cache; retries offline when the Hub is unreachable (as the server)."""
import os
from transformers import SuperPointForKeypointDetection
if os.path.isdir(str(name_or_path)):
model = SuperPointForKeypointDetection.from_pretrained(name_or_path)
else:
rev = revision if revision is not None else (DEFAULT_REVISION if name_or_path == DEFAULT_MODEL_ID else None)
try:
model = SuperPointForKeypointDetection.from_pretrained(name_or_path, revision=rev,
local_files_only=local_files_only)
except Exception as e: # noqa: BLE001 - network down, snapshot cached
if local_files_only or rev is None:
raise
LOG.warning("Hub resolution failed (%s: %s); retrying from the local HF cache", type(e).__name__, e)
model = SuperPointForKeypointDetection.from_pretrained(name_or_path, revision=rev, local_files_only=True)
return model.eval()
def _release(device, tt_model, tt_in, *, close: bool) -> None:
if device is None and tt_in is None:
if tt_model is not None:
tt_model.release()
return
import ttnn
try:
if device is not None:
ttnn.synchronize_device(device)
if tt_model is not None:
tt_model.release()
if tt_in is not None:
ttnn.deallocate(tt_in)
except Exception as e: # noqa: BLE001 - never mask the original error
LOG.warning("device tensor release failed: %s", e)
finally:
if close and device is not None:
ttnn.close_device(device)