superpoint-p150 / code /models /tt /postprocess.py
changh95's picture
Optimized build (2026-10-03): trace 3.73 -> 0.50 ms
c699c4c verified
Raw History Blame Contribute Delete
16.2 kB
# SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc.
# SPDX-License-Identifier: Apache-2.0
"""Host-side SuperPoint post-processing (torch only, no ttnn import).
This is the sequence the port validated in ``models/tests/test_superpoint.py``
(``_device_to_host_post`` -> ``_decode_keypoints(apply_nms=True)`` ->
``_extract_keypoints_single`` -> ``_sample_descriptors``), lifted out of the
``TtSuperPoint`` methods so that
* the serving app can call it with per-request ``nms_radius`` /
``keypoint_threshold`` / ``max_keypoints`` instead of the values frozen into
the model config, and
* it stays importable without ttnn (unit tests, tooling).
The device already applied the 65-way softmax (``ttnn.softmax`` in
``TtSuperPoint.run_device_compute``); nothing here applies it again. The input
``scores_nchw`` is the *softmaxed* score tensor exactly as
``superpoint_ttnn.device_outputs_to_host`` returns it.
"""
from __future__ import annotations
from typing import List, Tuple
import torch
import torch.nn.functional as F
DESCRIPTOR_SCALE = 8 # encoder stride: one descriptor cell per 8x8 pixels
def simple_nms(scores: torch.Tensor, nms_radius: int) -> torch.Tensor:
"""Single-pass NMS: keep pixels whose score equals the local (2r+1)^2 max.
The HF reference iterates a tie-expansion loop three times (~100 ms/iter on
host at 480x640). The single pass costs one max-pool and preserved
keypoint F1 98.8% @ top-500 / 2 px in the port's benchmark.
"""
if nms_radius <= 0:
return scores
pooled = F.max_pool2d(scores, kernel_size=nms_radius * 2 + 1, stride=1, padding=nms_radius)
return torch.where(scores == pooled, scores, torch.zeros_like(scores))
def fold_scores(scores_nchw: torch.Tensor, nms_radius: int | None) -> torch.Tensor:
"""(B, 65, h, w) softmaxed cell scores -> (B, 8h, 8w) dense map.
Drops the dustbin channel (64) and unfolds each 8x8 cell. ``nms_radius``
``None`` skips NMS (pre-NMS map); an int applies :func:`simple_nms`.
"""
scores = scores_nchw[:, :-1] # (B, 64, h, w)
b, _, fh, fw = scores.shape
scores = scores.permute(0, 2, 3, 1).reshape(b, fh, fw, 8, 8)
scores = scores.permute(0, 1, 3, 2, 4).reshape(b, fh * 8, fw * 8)
if nms_radius is not None:
scores = simple_nms(scores, nms_radius)
return scores
def extract_keypoints(
scores_1hw: torch.Tensor,
keypoint_threshold: float,
border_removal_distance: int,
max_keypoints: int,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Threshold, border-remove and top-k one (1, H, W) post-NMS map.
Returns ``keypoints`` as (N, 2) float ``(x, y)`` pixel coordinates in the
map's frame and ``scores`` as (N,). ``max_keypoints < 0`` keeps every point
above the threshold.
"""
_, height, width = scores_1hw.shape
keypoints = torch.nonzero(scores_1hw[0] > keypoint_threshold)
scores = scores_1hw[0][tuple(keypoints.t())]
border = border_removal_distance
mask_h = (keypoints[:, 0] >= border) & (keypoints[:, 0] < (height - border))
mask_w = (keypoints[:, 1] >= border) & (keypoints[:, 1] < (width - border))
mask = mask_h & mask_w
keypoints = keypoints[mask]
scores = scores[mask]
if max_keypoints >= 0 and keypoints.shape[0] > max_keypoints:
scores, idx = torch.topk(scores, max_keypoints, dim=0)
keypoints = keypoints[idx]
keypoints = torch.flip(keypoints, [1]).to(scores.dtype) # (y, x) -> (x, y)
return keypoints, scores
def sample_descriptors(
keypoints: torch.Tensor, descriptors: torch.Tensor, scale: int = DESCRIPTOR_SCALE
) -> torch.Tensor:
"""Bilinear-sample the (B, C, h, w) descriptor map at (B, N, 2) ``(x, y)`` points.
Returns (B, C, N), L2-normalised along C (the HF reference's
``_sample_descriptors``).
"""
batch_size, num_channels, height, width = descriptors.shape
keypoints = keypoints - scale / 2 + 0.5
divisor = torch.tensor([[(width * scale - scale / 2 - 0.5), (height * scale - scale / 2 - 0.5)]])
divisor = divisor.to(keypoints)
keypoints = keypoints / divisor
keypoints = keypoints * 2 - 1
keypoints = keypoints.view(batch_size, 1, -1, 2)
descriptors = F.grid_sample(descriptors, keypoints, mode="bilinear", align_corners=True)
descriptors = descriptors.reshape(batch_size, num_channels, -1)
descriptors = F.normalize(descriptors, p=2, dim=1)
return descriptors
def postprocess_keypoints(
scores_nchw: torch.Tensor,
descriptors_nchw: torch.Tensor,
*,
nms_radius: int,
keypoint_threshold: float,
max_keypoints: int,
border_removal_distance: int,
with_descriptors: bool = True,
) -> List[Tuple[torch.Tensor, torch.Tensor, torch.Tensor | None]]:
"""Full validated host post-processing for a batch.
``scores_nchw``: (B, 65, h, w) device-softmaxed scores as returned by
``superpoint_ttnn.device_outputs_to_host``. ``descriptors_nchw``:
(B, 256, h, w) device-L2-normalised descriptor map.
Returns one ``(keypoints (N, 2) xy, scores (N,), descriptors (N, 256) | None)``
triple per image, keypoints in the network-input pixel frame (480x640).
"""
scores_full = fold_scores(scores_nchw, nms_radius)
return postprocess_from_nms_map(
scores_full,
descriptors_nchw,
keypoint_threshold=keypoint_threshold,
max_keypoints=max_keypoints,
border_removal_distance=border_removal_distance,
with_descriptors=with_descriptors,
)
def postprocess_from_nms_map(
nms_map: torch.Tensor,
descriptors_nchw: torch.Tensor,
*,
keypoint_threshold: float,
max_keypoints: int,
border_removal_distance: int,
with_descriptors: bool = True,
) -> List[Tuple[torch.Tensor, torch.Tensor, torch.Tensor | None]]:
"""Post-processing from an already folded + NMS'd dense map (the ``TT_FUSED`` path).
``nms_map``: (B, H, W) post-NMS scores -- either ``fold_scores(scores_nchw, r)`` (host) or
the device NMS-T map ``TtSuperPoint.run_fused`` returns (bit-identical to it). Everything
after the NMS is the legacy code path: :func:`extract_keypoints` + :func:`sample_descriptors`.
"""
out = []
for i in range(nms_map.shape[0]):
kp, sc = extract_keypoints(
nms_map[i : i + 1], keypoint_threshold, border_removal_distance, max_keypoints
)
desc = None
if with_descriptors:
if kp.shape[0] > 0:
desc = sample_descriptors(kp[None], descriptors_nchw[i : i + 1])[0].transpose(0, 1)
else:
desc = torch.zeros((0, descriptors_nchw.shape[1]), dtype=descriptors_nchw.dtype)
out.append((kp, sc, desc))
return out
# ----------------------------------------------------------------------------- fast fused-path host side
# Same results as extract_keypoints / sample_descriptors on the fused outputs, without converting
# the full maps to fp32 NCHW: the NMS map stays bf16 and is thresholded on its bit pattern, and the
# descriptors are bilinearly sampled straight from the NHWC bf16 device readback (4 row gathers per
# keypoint instead of an NCHW permute + fp32 copy of the whole 4800x256 map).
def _bf16_threshold_bits(threshold: float) -> int:
"""Largest non-negative bf16 bit pattern p with value(p) <= float32(threshold). For x >= 0 in bf16:
float32(x) > float32(threshold) <=> bits(x) > p (bf16 patterns of non-negative values are
monotonic)."""
t = torch.tensor([threshold], dtype=torch.float32).view(torch.int32)
return int(t.item()) >> 16
def extract_keypoints_bf16(
nms_hw: torch.Tensor,
keypoint_threshold: float,
border_removal_distance: int,
max_keypoints: int,
) -> Tuple[torch.Tensor, torch.Tensor]:
""":func:`extract_keypoints` for one (H, W) **bf16** post-NMS map with scores >= 0 (the device
NMS output). Same keypoints, scores and order (raster order, then torch.topk on the same fp32
score vector), but the threshold and border mask run on a cropped int16 view of the bf16 map."""
if keypoint_threshold < 0 or nms_hw.dtype != torch.bfloat16:
return extract_keypoints(nms_hw.float()[None], keypoint_threshold, border_removal_distance, max_keypoints)
height, width = nms_hw.shape
b = border_removal_distance
inner = nms_hw[b : height - b, b : width - b]
mask = inner.view(torch.int16) > _bf16_threshold_bits(keypoint_threshold)
keypoints = torch.nonzero(mask)
scores = inner[mask].float()
if b:
keypoints += b
if max_keypoints >= 0 and keypoints.shape[0] > max_keypoints:
scores, idx = torch.topk(scores, max_keypoints, dim=0)
keypoints = keypoints[idx]
keypoints = torch.flip(keypoints, [1]).to(scores.dtype) # (y, x) -> (x, y)
return keypoints, scores
def bilinear_taps(keypoints: torch.Tensor, height: int, width: int, scale: int = DESCRIPTOR_SCALE):
"""The 4 bilinear taps of ``F.grid_sample(..., align_corners=True, padding zeros)`` at (N, 2)
``(x, y)`` keypoints on an (h, w) cell grid, with the same grid normalisation as
:func:`sample_descriptors`: list of 4 ``(flat cell index (clamped), weight * valid)`` pairs in
the order nw, ne, sw, se."""
kp = keypoints - scale / 2 + 0.5
divisor = torch.tensor([[(width * scale - scale / 2 - 0.5), (height * scale - scale / 2 - 0.5)]]).to(kp)
g = kp / divisor * 2 - 1
ix = ((g[:, 0] + 1) / 2) * (width - 1)
iy = ((g[:, 1] + 1) / 2) * (height - 1)
x0 = torch.floor(ix)
y0 = torch.floor(iy)
x1, y1 = x0 + 1, y0 + 1
w_nw = (x1 - ix) * (y1 - iy)
w_ne = (ix - x0) * (y1 - iy)
w_sw = (x1 - ix) * (iy - y0)
w_se = (ix - x0) * (iy - y0)
taps = []
for xx, yy, ww in ((x0, y0, w_nw), (x1, y0, w_ne), (x0, y1, w_sw), (x1, y1, w_se)):
valid = (xx >= 0) & (xx <= width - 1) & (yy >= 0) & (yy <= height - 1)
idx = (yy.clamp(0, height - 1) * width + xx.clamp(0, width - 1)).long()
taps.append((idx, ww * valid))
return taps
def sample_from_taps(taps, rows_of, channels: int) -> torch.Tensor:
"""Weighted sum of the 4 taps (fp32, nw+ne+sw+se in that order) then L2-normalise.
``rows_of(idx)`` returns the (N, C) descriptor rows of flat cell indices ``idx``."""
n = taps[0][0].shape[0]
out = torch.zeros((n, channels), dtype=torch.float32)
for idx, w in taps:
out += rows_of(idx).float() * w[:, None]
return F.normalize(out, p=2, dim=1)
def sample_descriptors_nhwc(
keypoints: torch.Tensor, descriptors_nhwc: torch.Tensor, scale: int = DESCRIPTOR_SCALE
) -> torch.Tensor:
""":func:`sample_descriptors` for one image from an NHWC (1, h, w, C) map (any float dtype) at
(N, 2) ``(x, y)`` points; returns (N, C) fp32, L2-normalised. Same grid normalisation and
bilinear / align_corners=True / zero-padding semantics as ``F.grid_sample`` (fp32 math), but only
the 4 neighbouring cells of each keypoint are gathered."""
_, height, width, channels = descriptors_nhwc.shape
flat = descriptors_nhwc.reshape(height * width, channels)
taps = bilinear_taps(keypoints, height, width, scale)
return sample_from_taps(taps, lambda idx: flat.index_select(0, idx), channels)
class SampleTables:
"""Exact per-axis factors of :func:`bilinear_taps` for every integer pixel coordinate.
``bilinear_taps`` is separable: the x part (x0, x1, x1-ix, ix-x0, validity) depends only on the
keypoint's x and the y part only on its y, and the weights are the fp32 products
(x-factor) * (y-factor). Evaluating the same elementwise fp32 expressions on arange(W) /
arange(H) once gives bit-identical factors; validity is folded in as a zero factor (the host
computes (wx*wy)*valid = 0 for an invalid tap; wx, wy >= 0 and finite, so 0*wy == that 0)."""
def __init__(self, height: int, width: int, scale: int = DESCRIPTOR_SCALE):
self.hc, self.wc = height // scale, width // scale
self.x0, self.x1, self.xl, self.xr = self._axis(width, self.wc, scale)
self.y0, self.y1, self.yt, self.yb = self._axis(height, self.hc, scale)
@staticmethod
def _axis(n_px: int, n_cells: int, scale: int):
v = torch.arange(n_px, dtype=torch.float32)
k = v - scale / 2 + 0.5
div = torch.tensor([(n_cells * scale - scale / 2 - 0.5)], dtype=torch.float32)
g = k / div * 2 - 1
i = ((g + 1) / 2) * (n_cells - 1)
i0 = torch.floor(i)
i1 = i0 + 1
lo = (i1 - i) * ((i0 >= 0) & (i0 <= n_cells - 1))
hi = (i - i0) * ((i1 >= 0) & (i1 <= n_cells - 1))
return i0.long().clamp(0, n_cells - 1), i1.long().clamp(0, n_cells - 1), lo, hi
def weight_table(self) -> torch.Tensor:
"""fp32 [H, 4*W]: the nw, ne, sw, se weights of every pixel (x factor * y factor, the
exact fp32 products :func:`bilinear_taps` computes for a keypoint at that pixel)."""
yt, yb = self.yt[:, None], self.yb[:, None]
xl, xr = self.xl[None, :], self.xr[None, :]
return torch.stack([xl * yt, xr * yt, xl * yb, xr * yb], -1).reshape(yt.shape[0], -1)
def header_upload(self, keypoints: torch.Tensor, slots: int) -> torch.Tensor:
"""(N, 2) integer-valued ``(x, y)`` keypoints -> int32 [1, 16 + 4*slots] in the device
keypoint-header format of ``kp_compact.cpp`` ([2] = N; per keypoint (y << 16) | x, score
(unused), (y0c*wc << 16) | y1c*wc, (x0c << 16) | x1c), the input of
``nms_kernels.DeviceSampler``'s second trace."""
n = keypoints.shape[0]
out = torch.zeros((1, 16 + 4 * slots), dtype=torch.int32)
out[0, 2] = n
kx, ky = keypoints[:, 0].long(), keypoints[:, 1].long()
e = out[0, 16 : 16 + 4 * n].view(n, 4)
e[:, 0] = ((ky << 16) | kx).to(torch.int32)
e[:, 2] = (((self.y0[ky] * self.wc) << 16) | (self.y1[ky] * self.wc)).to(torch.int32)
e[:, 3] = ((self.x0[kx] << 16) | self.x1[kx]).to(torch.int32)
return out
def decode_candidates(cand: torch.Tensor, slot_first_rows: torch.Tensor, width: int, max_keypoints: int):
"""Device keypoint-candidate slots (``nms_kernels.DeviceNms`` with a threshold; int32
[NSLOT, CAP+1]) -> ``(keypoints (N, 2) xy fp32, scores (N,) fp32)`` exactly as
:func:`extract_keypoints_bf16` returns them for the same NMS map (raster order, then
``torch.topk`` on the same fp32 score vector). Returns ``None`` when a slot overflowed."""
if cand.dtype != torch.int32:
cand = cand.view(torch.int32) if cand.element_size() == 4 else cand.to(torch.int32)
cap = cand.shape[1] - 1
counts = cand[:, 0]
if int(counts.max()) > cap:
return None
mask = torch.arange(cap)[None, :] < counts[:, None]
ev = cand[:, 1:][mask]
first = slot_first_rows[:, None].expand(-1, cap)[mask]
off = ev & 0xFFFF
y = first + torch.div(off, width, rounding_mode="floor")
x = off - torch.div(off, width, rounding_mode="floor") * width
scores = (ev >> 16).to(torch.int16).view(torch.bfloat16).float()
keypoints = torch.stack([y, x], 1).long()
if max_keypoints >= 0 and keypoints.shape[0] > max_keypoints:
scores, idx = torch.topk(scores, max_keypoints, dim=0)
keypoints = keypoints[idx]
keypoints = torch.flip(keypoints, [1]).to(scores.dtype)
return keypoints, scores
def postprocess_fused_bf16(
nms_map_bf16: torch.Tensor,
descriptors_nhwc: torch.Tensor | None,
*,
keypoint_threshold: float,
max_keypoints: int,
border_removal_distance: int,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor | None]:
"""One image: (H, W) bf16 device NMS map + (1, h, w, 256) NHWC descriptor readback ->
``(keypoints (N, 2) xy, scores (N,), descriptors (N, 256) | None)``; the same triple as
:func:`postprocess_from_nms_map` up to fp32 rounding of the bilinear sum."""
kp, sc = extract_keypoints_bf16(nms_map_bf16, keypoint_threshold, border_removal_distance, max_keypoints)
desc = None
if descriptors_nhwc is not None:
if kp.shape[0] > 0:
desc = sample_descriptors_nhwc(kp, descriptors_nhwc)
else:
desc = torch.zeros((0, descriptors_nhwc.shape[-1]), dtype=torch.float32)
return kp, sc, desc