1use crate::{DescriptorBuffer, DescriptorKind, FeatureMatch, VisionError, VisionResult};
4
5#[derive(Clone, Copy, Debug, Default, PartialEq)]
7pub struct MatchOptions {
8 pub cross_check: bool,
10 pub ratio: Option<f32>,
12 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
32pub 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}