Download code/models/tt/postprocess.py from changh95/superpoint-p150: direct link, hf CLI and curl.
- Browser
- Download file 16.2 kB
-
https://huggingface.co/changh95/superpoint-p150/resolve/main/code/models/tt/postprocess.py
- Command line
-
hf download hf://changh95/superpoint-p150/code/models/tt/postprocess.py
-
curl -L -o postprocess.py https://huggingface.co/changh95/superpoint-p150/resolve/main/code/models/tt/postprocess.py
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) | |
| 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 | |