//! TRACES: FR-MRG-8 //! The XFeat detector — the network under tract, and the decoder after it. //! //! Apache-2.0 weights (`models/LICENCE.md`), exported at a fixed shape by //! `tools/export-xfeat.sh` and loaded through the same `ort`-over-tract //! backend `dr-segment` and `dr-face` use, so this adds no runtime and no C //! to the tree. ~300 ms per frame on the reference desktop, ~400 ms on the //! tablet (S15.2, S15.4). use crate::features::{decode_xfeat, DecodeOptions, Features, XFeatMaps, DESCRIPTOR_LEN}; use crate::image::Gray; use crate::PanoError; /// The input shape the shipped export was made for. A different size is a /// different file (`tools/export-xfeat.sh`). pub const INPUT_WIDTH: usize = 1024; pub const INPUT_HEIGHT: usize = 768; #[cfg(feature = "embedded-model")] const EMBEDDED_MODEL: &[u8] = include_bytes!("../../../models/keypoints/xfeat-1024.onnx"); /// A loaded detector. pub struct XFeat { session: ort::session::Session, pub options: DecodeOptions, } impl XFeat { /// The weights compiled into the binary. #[cfg(feature = "embedded-model")] pub fn embedded() -> Result { Self::from_bytes(EMBEDDED_MODEL) } pub fn from_path(path: &std::path::Path) -> Result { let bytes = std::fs::read(path).map_err(PanoError::ModelRead)?; Self::from_bytes(&bytes) } pub fn from_bytes(bytes: &[u8]) -> Result { install_backend(); let session = ort::session::Session::builder() .map_err(PanoError::Inference)? .commit_from_memory(bytes) .map_err(PanoError::Inference)?; Ok(XFeat { session, options: DecodeOptions::default(), }) } /// Detect keypoints in an upright grayscale image. /// /// The image is fitted into the network's fixed input — scaled down if /// larger, never up, and padded to the right and bottom — and the /// keypoints come back in the coordinates of `image` itself, so a /// caller that already scaled a frame to a proxy maps them on with the /// scale it used and nothing else. pub fn detect(&mut self, image: &Gray) -> Result { let (fitted, scale) = image.fitted(INPUT_WIDTH, INPUT_HEIGHT); let padded = fitted.padded(INPUT_WIDTH, INPUT_HEIGHT); let input = ndarray::Array::from_shape_vec( ndarray::IxDyn(&[1, 1, INPUT_HEIGHT, INPUT_WIDTH]), padded.data, ) .expect("shape matches the buffer by construction"); let tensor = ort::value::Tensor::from_array(input).map_err(PanoError::Inference)?; let outputs = self .session .run(ort::inputs![tensor]) .map_err(PanoError::Inference)?; let (w8, h8) = (INPUT_WIDTH / 8, INPUT_HEIGHT / 8); let expect = |i: usize, channels: usize| -> Result, PanoError> { let (shape, data) = outputs[i] .try_extract_tensor::() .map_err(PanoError::Inference)?; let dims: Vec = shape.iter().copied().collect(); if dims != [1, channels as i64, h8 as i64, w8 as i64] { return Err(PanoError::Model(format!( "output {i} is {dims:?}, expected [1, {channels}, {h8}, {w8}] — \ not the export this decoder was written for" ))); } Ok(data.to_vec()) }; let feats = expect(0, DESCRIPTOR_LEN)?; let keypoints = expect(1, 65)?; let heatmap = expect(2, 1)?; let mut features = decode_xfeat( &XFeatMaps { feats: &feats, keypoints: &keypoints, heatmap: &heatmap, width: w8, height: h8, }, &self.options, ); // Back to the caller's image: drop anything the padding produced, // undo the fit. let border = self.options.border as f32; let limit_x = fitted.width as f32 - border; let limit_y = fitted.height as f32 - border; let mut kept_kp = Vec::with_capacity(features.len()); let mut kept_desc = Vec::with_capacity(features.descriptors.len()); for (i, kp) in features.keypoints.iter().enumerate() { if kp.x >= limit_x || kp.y >= limit_y { continue; } kept_kp.push(crate::features::Keypoint { x: (kp.x / scale as f32), y: (kp.y / scale as f32), score: kp.score, }); kept_desc.extend_from_slice(features.descriptor(i)); } features.keypoints = kept_kp; features.descriptors = kept_desc; features.width = image.width; features.height = image.height; Ok(features) } } fn install_backend() { use std::sync::Once; static ONCE: Once = Once::new(); ONCE.call_once(|| { // False if another crate installed it first, which is fine: there is // one backend compiled in for it to have chosen. let _ = ort::set_api(ort_tract::api()); }); }