Skip to main content

multivector/
fde.rs

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/// Scalar dot product. Kept as a portable fallback and used for tail elements
20/// past the SIMD-aligned prefix.
21#[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/// Public dot product with architecture dispatch. Routes to a NEON FMA
33/// implementation on aarch64 for the multiple-of-16 prefix and mops up the
34/// tail with the scalar path. Every hot path in the crate (FDE exhaustive
35/// scan, MaxSim rescoring) goes through this — reason enough to keep the
36/// dispatch cheap.
37#[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            // SAFETY: prefix is a multiple of 16 and prefix <= n <= left.len().
46            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/// NEON FP32 dot product for a length that is a multiple of 16.
57/// Uses four independent accumulators to hide FMA latency (~4 cycles on M-series).
58#[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/// Same as [`dot_neon_multiple_of_16`] but with the length fixed at 128.
90/// Const bound lets the compiler unroll the loop fully — in benchmarks this
91/// specialisation shaves ~30% off the general-dimension path (~2.4 ms vs
92/// ~3.2 ms on the 250-candidate MaxSim rescoring kernel).
93#[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
123/// ColBERT's late-interaction score: sum of per-query-token maxima.
124pub 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
141/// MaxSim over already-normalized query vectors and a flat document matrix.
142///
143/// Routes to a vectorised kernel when the dimension is a multiple of 16 on
144/// aarch64 (covers the standard ColBERT/E5/MPNet embedding sizes 128, 384,
145/// 512, 768, 1024). Falls back to the scalar path everywhere else.
146pub 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    // NOTE: query-outer / doc-inner order is intentional. A previous attempt
172    // to swap the loops (stream doc tokens once, iterate 32 query tokens
173    // per doc) regressed the kernel from 2.35 ms to 3.51 ms in the
174    // maxsim-bench harness. The original order keeps the current query
175    // vector pinned in registers across all 200 doc-token dot products,
176    // and lets the compiler track a single scalar `best` in a register
177    // across the inner loop. Doc-outer loses both benefits.
178    if dimension == 128 {
179        // Const-length specialisation for the standard ColBERT dim so the
180        // compiler can fully unroll the inner FMA loop.
181        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                // SAFETY: `dimension` is validated a multiple of 16 by the caller
205                // (maxsim_flat), q and doc are contiguous slices of `dimension`
206                // f32 values, so the NEON loads stay in-bounds.
207                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        // 100 is not a multiple of 16 — dispatch must take the scalar path.
261        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}