1pub type Vector = Vec<f32>;
2
3pub fn normalize(vector: &[f32]) -> Vector {
4 let norm = vector.iter().map(|x| x * x).sum::<f32>().sqrt();
5 if !norm.is_finite() {
6 let norm = vector
7 .iter()
8 .map(|&x| (x as f64).powi(2))
9 .sum::<f64>()
10 .sqrt();
11 vector.iter().map(|&x| (x as f64 / norm) as f32).collect()
12 } else if norm == 0.0 {
13 vec![0.0; vector.len()]
14 } else {
15 vector.iter().map(|x| x / norm).collect()
16 }
17}
18
19#[inline]
22fn dot_scalar(left: &[f32], right: &[f32]) -> f32 {
23 debug_assert_eq!(left.len(), right.len());
24 let n = left.len().min(right.len());
25 let mut s = 0.0f32;
26 for i in 0..n {
27 s += left[i] * right[i];
28 }
29 s
30}
31
32#[inline]
38pub fn dot(left: &[f32], right: &[f32]) -> f32 {
39 debug_assert_eq!(left.len(), right.len());
40 let n = left.len().min(right.len());
41 #[cfg(target_arch = "aarch64")]
42 {
43 let prefix = n & !15;
44 if prefix >= 16 {
45 let head = unsafe { dot_neon_multiple_of_16(left.as_ptr(), right.as_ptr(), prefix) };
47 if prefix == n {
48 return head;
49 }
50 return head + dot_scalar(&left[prefix..n], &right[prefix..n]);
51 }
52 }
53 dot_scalar(&left[..n], &right[..n])
54}
55
56#[cfg(target_arch = "aarch64")]
59#[inline(always)]
60unsafe fn dot_neon_multiple_of_16(a: *const f32, b: *const f32, len: usize) -> f32 {
61 unsafe {
62 use std::arch::aarch64::*;
63 debug_assert!(len % 16 == 0);
64 let mut acc0 = vdupq_n_f32(0.0);
65 let mut acc1 = vdupq_n_f32(0.0);
66 let mut acc2 = vdupq_n_f32(0.0);
67 let mut acc3 = vdupq_n_f32(0.0);
68 let mut i = 0usize;
69 while i < len {
70 let a0 = vld1q_f32(a.add(i));
71 let a1 = vld1q_f32(a.add(i + 4));
72 let a2 = vld1q_f32(a.add(i + 8));
73 let a3 = vld1q_f32(a.add(i + 12));
74 let b0 = vld1q_f32(b.add(i));
75 let b1 = vld1q_f32(b.add(i + 4));
76 let b2 = vld1q_f32(b.add(i + 8));
77 let b3 = vld1q_f32(b.add(i + 12));
78 acc0 = vfmaq_f32(acc0, a0, b0);
79 acc1 = vfmaq_f32(acc1, a1, b1);
80 acc2 = vfmaq_f32(acc2, a2, b2);
81 acc3 = vfmaq_f32(acc3, a3, b3);
82 i += 16;
83 }
84 let acc = vaddq_f32(vaddq_f32(acc0, acc1), vaddq_f32(acc2, acc3));
85 vaddvq_f32(acc)
86 }
87}
88
89#[cfg(target_arch = "aarch64")]
94#[inline(always)]
95unsafe fn dot_neon_128(a: *const f32, b: *const f32) -> f32 {
96 unsafe {
97 use std::arch::aarch64::*;
98 let mut acc0 = vdupq_n_f32(0.0);
99 let mut acc1 = vdupq_n_f32(0.0);
100 let mut acc2 = vdupq_n_f32(0.0);
101 let mut acc3 = vdupq_n_f32(0.0);
102 let mut i = 0usize;
103 while i < 128 {
104 let a0 = vld1q_f32(a.add(i));
105 let a1 = vld1q_f32(a.add(i + 4));
106 let a2 = vld1q_f32(a.add(i + 8));
107 let a3 = vld1q_f32(a.add(i + 12));
108 let b0 = vld1q_f32(b.add(i));
109 let b1 = vld1q_f32(b.add(i + 4));
110 let b2 = vld1q_f32(b.add(i + 8));
111 let b3 = vld1q_f32(b.add(i + 12));
112 acc0 = vfmaq_f32(acc0, a0, b0);
113 acc1 = vfmaq_f32(acc1, a1, b1);
114 acc2 = vfmaq_f32(acc2, a2, b2);
115 acc3 = vfmaq_f32(acc3, a3, b3);
116 i += 16;
117 }
118 let acc = vaddq_f32(vaddq_f32(acc0, acc1), vaddq_f32(acc2, acc3));
119 vaddvq_f32(acc)
120 }
121}
122
123pub fn maxsim(query: &[Vector], document: &[Vector]) -> f32 {
125 if query.is_empty() || document.is_empty() {
126 return 0.0;
127 }
128 let document: Vec<_> = document.iter().map(|v| normalize(v)).collect();
129 query
130 .iter()
131 .map(|query| {
132 let query = normalize(query);
133 document
134 .iter()
135 .map(|doc| dot(&query, doc))
136 .fold(f32::NEG_INFINITY, f32::max)
137 })
138 .sum()
139}
140
141pub fn maxsim_flat(query: &[Vector], document: &[f32], dimension: usize) -> f32 {
147 #[cfg(target_arch = "aarch64")]
148 {
149 if dimension > 0 && dimension % 16 == 0 && dimension <= 4096 {
150 return maxsim_flat_neon(query, document, dimension);
151 }
152 }
153 maxsim_flat_scalar(query, document, dimension)
154}
155
156#[inline]
157fn maxsim_flat_scalar(query: &[Vector], document: &[f32], dimension: usize) -> f32 {
158 query
159 .iter()
160 .map(|q| {
161 document
162 .chunks_exact(dimension)
163 .map(|d| dot(q, d))
164 .fold(f32::NEG_INFINITY, f32::max)
165 })
166 .sum()
167}
168
169#[cfg(target_arch = "aarch64")]
170fn maxsim_flat_neon(query: &[Vector], document: &[f32], dimension: usize) -> f32 {
171 if dimension == 128 {
179 return query
182 .iter()
183 .map(|q| {
184 debug_assert_eq!(q.len(), 128);
185 let qp = q.as_ptr();
186 let mut best = f32::NEG_INFINITY;
187 for doc in document.chunks_exact(128) {
188 let s = unsafe { dot_neon_128(qp, doc.as_ptr()) };
189 if s > best {
190 best = s;
191 }
192 }
193 best
194 })
195 .sum();
196 }
197 query
198 .iter()
199 .map(|q| {
200 debug_assert_eq!(q.len(), dimension);
201 let qp = q.as_ptr();
202 let mut best = f32::NEG_INFINITY;
203 for doc in document.chunks_exact(dimension) {
204 let s = unsafe { dot_neon_multiple_of_16(qp, doc.as_ptr(), dimension) };
208 if s > best {
209 best = s;
210 }
211 }
212 best
213 })
214 .sum()
215}
216
217#[cfg(test)]
218mod tests {
219 use super::*;
220
221 fn deterministic(seed: u64, dim: usize) -> Vector {
222 let mut s = seed.wrapping_mul(0x9E37_79B9_7F4A_7C15);
223 (0..dim)
224 .map(|_| {
225 s = s
226 .wrapping_mul(6364136223846793005)
227 .wrapping_add(1442695040888963407);
228 (((s >> 33) as u32) as f32 / u32::MAX as f32) * 2.0 - 1.0
229 })
230 .collect()
231 }
232
233 fn flat(doc: &[Vector]) -> (Vec<f32>, usize) {
234 let dim = doc[0].len();
235 (doc.iter().flat_map(|v| v.iter().copied()).collect(), dim)
236 }
237
238 #[test]
239 fn maxsim_flat_scalar_and_dispatch_agree_on_dim128() {
240 let query: Vec<_> = (0..8).map(|i| deterministic(0x11 + i, 128)).collect();
241 let doc_tokens: Vec<_> = (0..50).map(|i| deterministic(0x2000 + i, 128)).collect();
242 let (flat_doc, dim) = flat(&doc_tokens);
243 let s = maxsim_flat_scalar(&query, &flat_doc, dim);
244 let d = maxsim_flat(&query, &flat_doc, dim);
245 assert!((s - d).abs() < 1e-3, "scalar={s} dispatch={d}");
246 }
247
248 #[test]
249 fn maxsim_flat_scalar_and_dispatch_agree_on_dim384() {
250 let query: Vec<_> = (0..12).map(|i| deterministic(0x33 + i, 384)).collect();
251 let doc_tokens: Vec<_> = (0..30).map(|i| deterministic(0x4000 + i, 384)).collect();
252 let (flat_doc, dim) = flat(&doc_tokens);
253 let s = maxsim_flat_scalar(&query, &flat_doc, dim);
254 let d = maxsim_flat(&query, &flat_doc, dim);
255 assert!((s - d).abs() < 1e-3, "scalar={s} dispatch={d}");
256 }
257
258 #[test]
259 fn maxsim_flat_dispatch_falls_back_when_dim_not_multiple_of_16() {
260 let query: Vec<_> = (0..4).map(|i| deterministic(0x55 + i, 100)).collect();
262 let doc_tokens: Vec<_> = (0..10).map(|i| deterministic(0x6000 + i, 100)).collect();
263 let (flat_doc, dim) = flat(&doc_tokens);
264 let s = maxsim_flat_scalar(&query, &flat_doc, dim);
265 let d = maxsim_flat(&query, &flat_doc, dim);
266 assert!((s - d).abs() < 1e-6, "scalar={s} dispatch={d}");
267 }
268}