# 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