Files
DarkRoom/core/dr-pano/src/features.rs
T
dtourolle 231b4a54ab dr-pano: the geometry, from features to cameras
A new crate holding the CPU half of a merge (FR-MRG-10): the grayscale
proxy with orientation, the XFeat decoder ported step for step from the
reference detectAndCompute, mutual-nearest-neighbour matching, a robust
pairwise homography with the focal length read off it, a hand-rolled
Levenberg–Marquardt bundle adjustment over every rotation and the focal,
the three output projections, and align(), which chains it all and names
the frames it could not place rather than guessing (FR-MRG-5).

Dependency-free without the xfeat feature — linalg.rs says why the dense
algebra is hand-rolled — and tested on synthetic sweeps whose answer is
known exactly. The noise test records the single-row degeneracy: one
pixel of noise is a tenth of a percent of focal, which is a uniform
stretch of the sweep, not a misalignment.
2026-09-19 15:24:12 +02:00

337 lines
12 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! Keypoints with descriptors, and the decoder that reads them out of
//! XFeat's dense maps.
//!
//! The network (S15.2) produces three maps at an eighth of the input
//! resolution and stops; everything from there to a list of keypoints is
//! this file, in plain Rust, for the reason `dr-segment` decodes yolo26's
//! heads itself: the post-processing is cheap, shape-dependent and exactly
//! the kind of graph tract parses badly. It is a port of the reference
//! `XFeat.detectAndCompute`, step for step, so that a keypoint here is the
//! keypoint the paper's numbers were measured on.
/// One detected point, in the pixel coordinates of the image it was
/// detected in, with the detector's confidence.
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct Keypoint {
pub x: f32,
pub y: f32,
/// The reliability the detector assigned; higher is better, and the
/// scale is the detector's own — comparable within one model only.
pub score: f32,
}
/// The keypoints of one image and their descriptors.
#[derive(Debug, Clone, PartialEq)]
pub struct Features {
pub keypoints: Vec<Keypoint>,
/// `keypoints.len() × DESCRIPTOR_LEN`, each row L2-normalised, so that a
/// dot product between two rows is their cosine similarity.
pub descriptors: Vec<f32>,
/// The image the coordinates are in.
pub width: usize,
pub height: usize,
}
/// The length of one descriptor. XFeat's is 64; the matcher does not care
/// what the number is, only that both sides agree.
pub const DESCRIPTOR_LEN: usize = 64;
impl Features {
pub fn len(&self) -> usize {
self.keypoints.len()
}
pub fn is_empty(&self) -> bool {
self.keypoints.is_empty()
}
pub fn descriptor(&self, i: usize) -> &[f32] {
&self.descriptors[i * DESCRIPTOR_LEN..(i + 1) * DESCRIPTOR_LEN]
}
}
/// XFeat's three output maps, as the network hands them back.
///
/// All three are `channels × height × width` at an eighth of the input, in
/// the NCHW order the ONNX export declares (`feats [1, 64, H/8, W/8]`,
/// `keypoints [1, 65, H/8, W/8]`, `heatmap [1, 1, H/8, W/8]`).
pub struct XFeatMaps<'a> {
/// 64 channels: the dense descriptor field.
pub feats: &'a [f32],
/// 65 channels: for each 8×8 cell, a logit per position plus one for
/// "no keypoint here".
pub keypoints: &'a [f32],
/// 1 channel: reliability.
pub heatmap: &'a [f32],
/// The maps' width and height (the input's, divided by eight).
pub width: usize,
pub height: usize,
}
/// How the decoder picks keypoints.
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct DecodeOptions {
/// Keep at most this many, by score. The reference default is 4096.
pub top_k: usize,
/// A cell position's softmax probability must exceed this to be a
/// keypoint at all. The reference default is 0.05.
pub threshold: f32,
/// Ignore keypoints within this many pixels of the map's edge. A frame
/// padded into the detector's fixed input (`Gray::padded`) has a hard
/// edge where the padding starts, and the detector fires on it.
pub border: usize,
}
impl Default for DecodeOptions {
fn default() -> Self {
DecodeOptions {
top_k: 4096,
threshold: 0.05,
border: 4,
}
}
}
/// Decode keypoints and descriptors from the network's maps.
///
/// The reference, step for step:
/// 1. softmax over the 65 logits of each cell, keep the 64 positions;
/// 2. pixel-shuffle those into a full-resolution keypoint heatmap — channel
/// `c` of cell `(cx, cy)` is pixel `(cx·8 + c%8, cy·8 + c/8)`;
/// 3. 5×5 non-maximum suppression over that heatmap, above `threshold`;
/// 4. score each survivor by its heatmap value times the reliability map
/// sampled bilinearly at its position;
/// 5. keep the `top_k` by score;
/// 6. sample the descriptor field bilinearly at each and L2-normalise.
///
/// Bilinear where the reference samples the descriptor field bicubically:
/// a quarter-pixel's difference in a field that is smooth by construction,
/// and one interpolator rather than two to keep correct.
pub fn decode_xfeat(maps: &XFeatMaps<'_>, opts: &DecodeOptions) -> Features {
let (w8, h8) = (maps.width, maps.height);
let (w, h) = (w8 * 8, h8 * 8);
let cells = w8 * h8;
debug_assert_eq!(maps.keypoints.len(), 65 * cells);
debug_assert_eq!(maps.feats.len(), DESCRIPTOR_LEN * cells);
debug_assert_eq!(maps.heatmap.len(), cells);
// 1 + 2: softmax per cell, scattered into the full-resolution heatmap.
let mut heat = vec![0.0f32; w * h];
for cy in 0..h8 {
for cx in 0..w8 {
let cell = cy * w8 + cx;
let logit = |c: usize| maps.keypoints[c * cells + cell];
let max = (0..65).map(logit).fold(f32::MIN, f32::max);
let mut sum = 0.0f32;
let mut exps = [0.0f32; 65];
for (c, e) in exps.iter_mut().enumerate() {
*e = (logit(c) - max).exp();
sum += *e;
}
for (c, e) in exps.iter().enumerate().take(64) {
let (dx, dy) = (c % 8, c / 8);
heat[(cy * 8 + dy) * w + cx * 8 + dx] = e / sum;
}
}
}
// 3: a pixel survives if it is the maximum of its 5×5 neighbourhood and
// above threshold. Ties go to every tied pixel, as the reference's
// `x == max_pool(x)` does.
let border = opts.border.max(2);
let mut survivors: Vec<(usize, usize, f32)> = Vec::new();
for y in border..h.saturating_sub(border) {
for x in border..w.saturating_sub(border) {
let v = heat[y * w + x];
if v <= opts.threshold {
continue;
}
let mut is_max = true;
'nb: for ny in y - 2..=y + 2 {
for nx in x - 2..=x + 2 {
if heat[ny * w + nx] > v {
is_max = false;
break 'nb;
}
}
}
if is_max {
survivors.push((x, y, v));
}
}
}
// 4: heatmap value × reliability, the latter sampled at the keypoint's
// position in map coordinates (`align_corners = False`: pixel `x` of the
// full image is `x / 8 - 0.5` in the map).
let sample = |field: &[f32], channels: usize, c: usize, x: f32, y: f32| -> f32 {
let fx = (x / 8.0 - 0.5).clamp(0.0, (w8 - 1) as f32);
let fy = (y / 8.0 - 0.5).clamp(0.0, (h8 - 1) as f32);
let x0 = fx as usize;
let y0 = fy as usize;
let x1 = (x0 + 1).min(w8 - 1);
let y1 = (y0 + 1).min(h8 - 1);
let tx = fx - x0 as f32;
let ty = fy - y0 as f32;
let at = |xx: usize, yy: usize| field[c * (w8 * h8) + yy * w8 + xx];
let _ = channels;
let top = at(x0, y0) * (1.0 - tx) + at(x1, y0) * tx;
let bot = at(x0, y1) * (1.0 - tx) + at(x1, y1) * tx;
top * (1.0 - ty) + bot * ty
};
let mut scored: Vec<(usize, usize, f32)> = survivors
.into_iter()
.map(|(x, y, v)| {
let r = sample(maps.heatmap, 1, 0, x as f32, y as f32);
(x, y, v * r)
})
.collect();
// 5: best first, then cut. `sort_unstable_by` on a total order of the
// score; NaN cannot occur — every input is a probability or a sigmoid.
scored.sort_unstable_by(|a, b| b.2.total_cmp(&a.2));
scored.truncate(opts.top_k);
// 6: descriptors.
let mut keypoints = Vec::with_capacity(scored.len());
let mut descriptors = Vec::with_capacity(scored.len() * DESCRIPTOR_LEN);
for (x, y, score) in scored {
let (xf, yf) = (x as f32, y as f32);
let start = descriptors.len();
for c in 0..DESCRIPTOR_LEN {
descriptors.push(sample(maps.feats, DESCRIPTOR_LEN, c, xf, yf));
}
let norm = descriptors[start..]
.iter()
.map(|v| v * v)
.sum::<f32>()
.sqrt()
.max(1e-12);
for v in &mut descriptors[start..] {
*v /= norm;
}
keypoints.push(Keypoint {
x: xf,
y: yf,
score,
});
}
Features {
keypoints,
descriptors,
width: w,
height: h,
}
}
#[cfg(test)]
mod tests {
use super::*;
/// Maps for a `w8 × h8` grid where every cell says "no keypoint" except
/// the listed ones, which put all their weight on one position.
fn maps(w8: usize, h8: usize, hot: &[(usize, usize, usize)]) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
let cells = w8 * h8;
let mut kp = vec![0.0f32; 65 * cells];
// "None" strongly preferred everywhere.
for cell in 0..cells {
kp[64 * cells + cell] = 10.0;
}
for &(cx, cy, c) in hot {
let cell = cy * w8 + cx;
kp[64 * cells + cell] = 0.0;
kp[c * cells + cell] = 10.0;
}
let heat = vec![0.5f32; cells];
// Descriptors: channel c is constant c across the field, so any
// sampled descriptor is the same known vector.
let mut feats = vec![0.0f32; DESCRIPTOR_LEN * cells];
for c in 0..DESCRIPTOR_LEN {
for v in &mut feats[c * cells..(c + 1) * cells] {
*v = c as f32;
}
}
(feats, kp, heat)
}
#[test]
fn a_hot_cell_position_becomes_a_keypoint_at_the_right_pixel() {
// Cell (2, 1), channel 8*3 + 5 = 29 → pixel (2*8 + 5, 1*8 + 3).
let (f, k, h) = maps(8, 8, &[(2, 1, 29)]);
let out = decode_xfeat(
&XFeatMaps {
feats: &f,
keypoints: &k,
heatmap: &h,
width: 8,
height: 8,
},
&DecodeOptions::default(),
);
assert_eq!(out.len(), 1);
assert_eq!((out.keypoints[0].x, out.keypoints[0].y), (21.0, 11.0));
assert_eq!((out.width, out.height), (64, 64));
// Score is the softmax weight (~1) times the reliability (0.5).
assert!((out.keypoints[0].score - 0.5).abs() < 5e-3);
}
#[test]
fn descriptors_are_unit_length() {
let (f, k, h) = maps(8, 8, &[(3, 3, 0), (5, 5, 63)]);
let out = decode_xfeat(
&XFeatMaps {
feats: &f,
keypoints: &k,
heatmap: &h,
width: 8,
height: 8,
},
&DecodeOptions::default(),
);
assert_eq!(out.len(), 2);
for i in 0..2 {
let n: f32 = out.descriptor(i).iter().map(|v| v * v).sum();
assert!((n - 1.0).abs() < 1e-5);
}
}
#[test]
fn top_k_keeps_the_best() {
let (f, k, mut h) = maps(8, 8, &[(1, 1, 0), (3, 3, 0), (5, 5, 0)]);
// Make cell (3, 3) the most reliable.
h[3 * 8 + 3] = 0.9;
let out = decode_xfeat(
&XFeatMaps {
feats: &f,
keypoints: &k,
heatmap: &h,
width: 8,
height: 8,
},
&DecodeOptions {
top_k: 1,
..Default::default()
},
);
assert_eq!(out.len(), 1);
assert_eq!((out.keypoints[0].x, out.keypoints[0].y), (24.0, 24.0));
}
#[test]
fn the_border_is_excluded() {
let (f, k, h) = maps(8, 8, &[(0, 0, 0)]);
let out = decode_xfeat(
&XFeatMaps {
feats: &f,
keypoints: &k,
heatmap: &h,
width: 8,
height: 8,
},
&DecodeOptions::default(),
);
assert!(out.is_empty());
}
}