1use crate::types::Multivector;
21
22fn rows(multivector: &Multivector) -> Vec<Vec<f32>> {
24 multivector.to_f32()
25}
26
27fn score_against(query: &[Vec<f32>], document: &[Vec<f32>]) -> f32 {
28 query
29 .iter()
30 .map(|query_token| {
31 document
32 .iter()
33 .map(|document_token| dot(query_token, document_token))
34 .fold(f32::NEG_INFINITY, f32::max)
35 })
36 .map(|best| if best.is_finite() { best } else { 0.0 })
38 .sum()
39}
40
41fn dot(left: &[f32], right: &[f32]) -> f32 {
42 left.iter().zip(right).map(|(a, b)| a * b).sum()
43}
44
45pub fn maxsim(query: &Multivector, documents: &[Multivector]) -> Vec<f32> {
49 let query = rows(query);
50 documents
51 .iter()
52 .map(|document| score_against(&query, &rows(document)))
54 .collect()
55}
56
57pub fn maxsim_batch(queries: &[Multivector], documents: &[Multivector]) -> Vec<Vec<f32>> {
62 let queries: Vec<Vec<Vec<f32>>> = queries.iter().map(rows).collect();
63 let mut scores = vec![vec![0.0f32; documents.len()]; queries.len()];
64
65 for (document_index, document) in documents.iter().enumerate() {
68 let document = rows(document);
69 for (query_index, query) in queries.iter().enumerate() {
70 scores[query_index][document_index] = score_against(query, &document);
71 }
72 }
73 scores
74}
75
76#[cfg(test)]
77mod tests {
78 #![allow(clippy::float_cmp)]
80
81 use super::*;
82 use half::f16;
83
84 fn f32_mv(rows: &[&[f32]]) -> Multivector {
85 Multivector::F32(rows.iter().map(|row| row.to_vec()).collect())
86 }
87
88 fn f16_mv(rows: &[&[f32]]) -> Multivector {
89 Multivector::F16(
90 rows.iter()
91 .map(|row| row.iter().map(|v| f16::from_f32(*v)).collect())
92 .collect(),
93 )
94 }
95
96 #[test]
97 fn identical_orthonormal_matrices_score_the_query_length() {
98 let query = f32_mv(&[&[1.0, 0.0], &[0.0, 1.0]]);
99 let scores = maxsim(&query, std::slice::from_ref(&query));
100 assert_eq!(scores.len(), 1);
101 approx::assert_relative_eq!(scores[0], 2.0, epsilon = 1e-5);
102 }
103
104 #[test]
105 fn scores_rank_by_similarity() {
106 let query = f32_mv(&[&[1.0, 0.0]]);
107 let documents = [
108 f32_mv(&[&[1.0, 0.0]]), f32_mv(&[&[std::f32::consts::FRAC_1_SQRT_2; 2]]), f32_mv(&[&[0.0, 1.0]]), ];
112 let scores = maxsim(&query, &documents);
113 assert!(scores[0] > scores[1] && scores[1] > scores[2]);
114 approx::assert_relative_eq!(scores[2], 0.0, epsilon = 1e-6);
115 }
116
117 #[test]
118 fn max_is_over_document_tokens_and_sum_is_over_query_tokens() {
119 let query = f32_mv(&[&[1.0, 0.0], &[0.0, 1.0]]);
121 let document = f32_mv(&[&[0.0, 1.0], &[1.0, 0.0], &[0.0, 0.0]]);
122 approx::assert_relative_eq!(maxsim(&query, &[document])[0], 2.0, epsilon = 1e-6);
123 }
124
125 #[test]
126 fn f16_inputs_score_exactly_as_their_widened_selves() {
127 let query = f16_mv(&[&[0.3, -0.7], &[0.1, 0.9]]);
128 let documents = [f16_mv(&[&[0.2, 0.5], &[-0.4, 0.8]]), f16_mv(&[&[1.0, 0.0]])];
129
130 let widened_query = Multivector::F32(query.to_f32());
131 let widened_documents: Vec<Multivector> = documents
132 .iter()
133 .map(|d| Multivector::F32(d.to_f32()))
134 .collect();
135
136 assert_eq!(
139 maxsim(&query, &documents),
140 maxsim(&widened_query, &widened_documents)
141 );
142 }
143
144 #[test]
145 fn batch_agrees_with_the_single_query_form() {
146 let queries = [f32_mv(&[&[1.0, 0.0]]), f32_mv(&[&[0.0, 1.0], &[1.0, 0.0]])];
147 let documents = [f32_mv(&[&[1.0, 0.0]]), f32_mv(&[&[0.6, 0.8]])];
148
149 let batch = maxsim_batch(&queries, &documents);
150 assert_eq!(batch.len(), 2);
151 for (index, query) in queries.iter().enumerate() {
152 let single = maxsim(query, &documents);
153 for (document_index, score) in single.iter().enumerate() {
154 approx::assert_relative_eq!(batch[index][document_index], score, epsilon = 1e-6);
155 }
156 }
157 }
158
159 #[test]
160 fn variable_token_counts_all_produce_finite_scores() {
161 let query = f32_mv(&[&[1.0, 0.0]]);
162 for count in [1usize, 10, 100] {
163 let document = Multivector::F32(
164 (0..count)
165 .map(|i| vec![i as f32 / count as f32, 0.5])
166 .collect(),
167 );
168 assert!(maxsim(&query, &[document])[0].is_finite());
169 }
170 }
171
172 #[test]
173 fn empty_inputs_are_handled_rather_than_producing_infinities() {
174 let query = f32_mv(&[&[1.0, 0.0]]);
175 assert!(maxsim(&query, &[]).is_empty());
176 assert_eq!(maxsim(&query, &[Multivector::F32(Vec::new())]), vec![0.0]);
177 assert_eq!(maxsim(&Multivector::F32(Vec::new()), &[query]), vec![0.0]);
178 assert!(maxsim_batch(&[], &[]).is_empty());
179 }
180}