Skip to main content

spatialrust_vision/
matcher.rs

1//! Brute-force descriptor matching with explicit distance semantics.
2
3use crate::{DescriptorBuffer, DescriptorKind, FeatureMatch, VisionError, VisionResult};
4
5/// Filtering applied to brute-force descriptor correspondences.
6#[derive(Clone, Copy, Debug, Default, PartialEq)]
7pub struct MatchOptions {
8    /// Keep only pairs whose reverse nearest neighbour is the query row.
9    pub cross_check: bool,
10    /// Lowe ratio threshold in `(0, 1)`; requires at least two train rows.
11    pub ratio: Option<f32>,
12    /// Optional inclusive maximum distance.
13    pub max_distance: Option<f32>,
14}
15
16impl MatchOptions {
17    fn validate(self) -> VisionResult<Self> {
18        if self.ratio.is_some_and(|ratio| !ratio.is_finite() || ratio <= 0.0 || ratio >= 1.0) {
19            return Err(VisionError::InvalidParameter(
20                "descriptor match ratio must be finite and in (0, 1)".into(),
21            ));
22        }
23        if self.max_distance.is_some_and(|distance| !distance.is_finite() || distance < 0.0) {
24            return Err(VisionError::InvalidParameter(
25                "descriptor maximum distance must be finite and non-negative".into(),
26            ));
27        }
28        Ok(self)
29    }
30}
31
32/// Matches each query descriptor to its nearest train descriptor.
33///
34/// Binary rows use Hamming distance and float rows use Euclidean L2 distance.
35/// Equal distances are resolved by the lowest row index. Returned matches remain
36/// in ascending query-row order.
37pub fn match_descriptors(
38    query: &DescriptorBuffer,
39    train: &DescriptorBuffer,
40    options: MatchOptions,
41) -> VisionResult<Vec<FeatureMatch>> {
42    let options = options.validate()?;
43    validate_compatibility(query, train)?;
44    if query.is_empty() || train.is_empty() {
45        return Ok(Vec::new());
46    }
47    if options.ratio.is_some() && train.len() < 2 {
48        return Err(VisionError::InvalidParameter(
49            "descriptor ratio matching requires at least two train rows".into(),
50        ));
51    }
52
53    let reverse_best = options
54        .cross_check
55        .then(|| (0..train.len()).map(|index| nearest(train, index, query).0).collect::<Vec<_>>());
56    let mut matches = Vec::with_capacity(query.len());
57    for query_index in 0..query.len() {
58        let (train_index, best, second) = nearest(query, query_index, train);
59        if options.ratio.is_some_and(|ratio| best >= ratio * second.unwrap_or(f32::INFINITY)) {
60            continue;
61        }
62        if options.max_distance.is_some_and(|maximum| best > maximum) {
63            continue;
64        }
65        if reverse_best.as_ref().is_some_and(|indices| indices[train_index] != query_index) {
66            continue;
67        }
68        matches.push(FeatureMatch::try_new(query_index, train_index, best)?);
69    }
70    Ok(matches)
71}
72
73fn validate_compatibility(query: &DescriptorBuffer, train: &DescriptorBuffer) -> VisionResult<()> {
74    if query.kind() != train.kind() || query.width() != train.width() {
75        return Err(VisionError::ShapeMismatch(format!(
76            "descriptor matrices must have equal kind and width (query {:?}/{}, train {:?}/{})",
77            query.kind(),
78            query.width(),
79            train.kind(),
80            train.width()
81        )));
82    }
83    Ok(())
84}
85
86fn nearest(
87    source: &DescriptorBuffer,
88    source_index: usize,
89    target: &DescriptorBuffer,
90) -> (usize, f32, Option<f32>) {
91    let mut candidates = (0..target.len())
92        .map(|target_index| (target_index, distance(source, source_index, target, target_index)))
93        .collect::<Vec<_>>();
94    candidates.sort_by(|left, right| left.1.total_cmp(&right.1).then_with(|| left.0.cmp(&right.0)));
95    let (best_index, best_distance) = candidates[0];
96    (best_index, best_distance, candidates.get(1).map(|candidate| candidate.1))
97}
98
99fn distance(
100    left: &DescriptorBuffer,
101    left_index: usize,
102    right: &DescriptorBuffer,
103    right_index: usize,
104) -> f32 {
105    match left.kind() {
106        DescriptorKind::Binary => left
107            .binary_row(left_index)
108            .expect("validated binary row")
109            .iter()
110            .zip(right.binary_row(right_index).expect("validated binary row"))
111            .map(|(a, b)| (a ^ b).count_ones())
112            .sum::<u32>() as f32,
113        DescriptorKind::Float32 => left
114            .float32_row(left_index)
115            .expect("validated float row")
116            .iter()
117            .zip(right.float32_row(right_index).expect("validated float row"))
118            .map(|(a, b)| {
119                let delta = a - b;
120                delta * delta
121            })
122            .sum::<f32>()
123            .sqrt(),
124    }
125}
126
127#[cfg(test)]
128mod tests {
129    use super::{match_descriptors, MatchOptions};
130    use crate::DescriptorBuffer;
131
132    #[test]
133    fn hamming_matching_is_deterministic_and_filters_distance() {
134        let query = DescriptorBuffer::try_binary(2, 1, vec![0b0000_0000, 0b1111_0000]).unwrap();
135        let train = DescriptorBuffer::try_binary(3, 1, vec![0b0000_0011, 0b0000_1100, 0b1111_1111])
136            .unwrap();
137        let matches = match_descriptors(&query, &train, MatchOptions::default()).unwrap();
138        assert_eq!((matches[0].train_index(), matches[0].distance()), (0, 2.0));
139        assert_eq!((matches[1].train_index(), matches[1].distance()), (2, 4.0));
140
141        let filtered = match_descriptors(
142            &query,
143            &train,
144            MatchOptions { max_distance: Some(2.0), ..MatchOptions::default() },
145        )
146        .unwrap();
147        assert_eq!(filtered.len(), 1);
148    }
149
150    #[test]
151    fn l2_ratio_and_cross_check_match_expected_rows() {
152        let query = DescriptorBuffer::try_float32(2, 2, vec![0.0, 0.0, 10.0, 10.0]).unwrap();
153        let train =
154            DescriptorBuffer::try_float32(3, 2, vec![1.0, 0.0, 3.0, 0.0, 10.0, 9.0]).unwrap();
155        let matches = match_descriptors(
156            &query,
157            &train,
158            MatchOptions { cross_check: true, ratio: Some(0.8), max_distance: None },
159        )
160        .unwrap();
161        assert_eq!(matches.len(), 2);
162        assert_eq!((matches[0].query_index(), matches[0].train_index()), (0, 0));
163        assert_eq!((matches[1].query_index(), matches[1].train_index()), (1, 2));
164        assert_eq!(matches[0].distance(), 1.0);
165        assert_eq!(matches[1].distance(), 1.0);
166    }
167
168    #[test]
169    fn incompatible_and_invalid_match_options_are_rejected() {
170        let binary = DescriptorBuffer::try_binary(1, 1, vec![0]).unwrap();
171        let float = DescriptorBuffer::try_float32(1, 1, vec![0.0]).unwrap();
172        assert!(match_descriptors(&binary, &float, MatchOptions::default()).is_err());
173        assert!(match_descriptors(
174            &binary,
175            &binary,
176            MatchOptions { ratio: Some(1.0), ..MatchOptions::default() }
177        )
178        .is_err());
179    }
180}