//! 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, /// `keypoints.len() × DESCRIPTOR_LEN`, each row L2-normalised, so that a /// dot product between two rows is their cosine similarity. pub descriptors: Vec, /// 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::() .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, Vec, Vec) { 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()); } }