Skip to main content

multivector/
fde.rs

1#[cfg(target_arch = "x86_64")]
2use annex::vector::simd::{CpuLevel, cpu_level};
3
4pub type Vector = Vec<f32>;
5
6/// Whether [`dot`] and [`maxsim_flat`] are using the AVX-512 BF16 kernels. That needs a CPU
7/// with the instructions and an explicit opt-in (`VECTORDB_BF16=1`); see
8/// `annex::vector::simd::bf16_enabled` for why it is not automatic.
9pub(crate) fn bf16_active() -> bool {
10    #[cfg(target_arch = "x86_64")]
11    {
12        matches!(cpu_level(), CpuLevel::Avx512Bf16)
13    }
14    #[cfg(not(target_arch = "x86_64"))]
15    {
16        false
17    }
18}
19
20/// Worst-case error of a BF16 dot product of two unit vectors.
21///
22/// The BF16 kernels round each operand to bfloat16 (8 significant bits, so a relative error
23/// of at most 2^-8) and then multiply-accumulate in f32. Each product therefore has a
24/// relative error of at most 2 * 2^-8 (plus a second-order 2^-16), which gives
25/// `|bf16_dot - dot| <= (2^-7 + 2^-16) * sum|a_i * b_i|`. For unit vectors
26/// `sum|a_i * b_i| <= |a| * |b| = 1` (Cauchy-Schwarz), so the bound is about 7.8e-3 and does
27/// not grow with the dimension. The constant adds slack for f32 accumulation.
28///
29/// A tolerance of 1e-3 is not attainable in BF16: the self-dot of a normalised
30/// `1/(i+1)` vector is off by 3.9e-3 at dim 64 (see
31/// `bf16_operand_rounding_stays_within_the_documented_bound`).
32#[cfg(test)]
33pub(crate) const BF16_UNIT_DOT_TOL: f32 = 8.2e-3;
34
35pub fn normalize(vector: &[f32]) -> Vector {
36    let norm = vector.iter().map(|x| x * x).sum::<f32>().sqrt();
37    if !norm.is_finite() {
38        let norm = vector
39            .iter()
40            .map(|&x| (x as f64).powi(2))
41            .sum::<f64>()
42            .sqrt();
43        vector.iter().map(|&x| (x as f64 / norm) as f32).collect()
44    } else if norm == 0.0 {
45        vec![0.0; vector.len()]
46    } else {
47        vector.iter().map(|x| x / norm).collect()
48    }
49}
50
51/// Scalar dot product. Kept as a portable fallback and used for tail elements
52/// past the SIMD-aligned prefix.
53#[cfg(target_arch = "aarch64")]
54#[inline]
55fn dot_scalar(left: &[f32], right: &[f32]) -> f32 {
56    debug_assert_eq!(left.len(), right.len());
57    let n = left.len().min(right.len());
58    let mut s = 0.0f32;
59    for i in 0..n {
60        s += left[i] * right[i];
61    }
62    s
63}
64
65/// Public dot product with architecture dispatch. Routes to a NEON FMA
66/// implementation on aarch64 for the multiple-of-16 prefix and mops up the
67/// tail with the scalar path; x86_64 uses `annex::vector::kernels::dot`. Every hot path in the crate (FDE exhaustive
68/// scan, MaxSim rescoring) goes through this — reason enough to keep the
69/// dispatch cheap.
70#[inline]
71pub fn dot(left: &[f32], right: &[f32]) -> f32 {
72    debug_assert_eq!(left.len(), right.len());
73    let n = left.len().min(right.len());
74    #[cfg(target_arch = "aarch64")]
75    {
76        let prefix = n & !15;
77        if prefix >= 16 {
78            // SAFETY: prefix is a multiple of 16 and prefix <= n <= left.len().
79            let head = unsafe { dot_neon_multiple_of_16(left.as_ptr(), right.as_ptr(), prefix) };
80            if prefix == n {
81                return head;
82            }
83            return head + dot_scalar(&left[prefix..n], &right[prefix..n]);
84        }
85        dot_scalar(&left[..n], &right[..n])
86    }
87    #[cfg(target_arch = "x86_64")]
88    if n > 0 && bf16_active() {
89        return unsafe { dot_avx512_bf16_len(left.as_ptr(), right.as_ptr(), n) };
90    }
91    #[cfg(not(target_arch = "aarch64"))]
92    {
93        // AVX-512 / AVX2+FMA with runtime dispatch; scalar elsewhere.
94        annex::vector::kernels::dot(&left[..n], &right[..n])
95    }
96}
97
98/// NEON FP32 dot product for a length that is a multiple of 16.
99/// Uses four independent accumulators to hide FMA latency (~4 cycles on M-series).
100#[cfg(target_arch = "aarch64")]
101#[inline(always)]
102unsafe fn dot_neon_multiple_of_16(a: *const f32, b: *const f32, len: usize) -> f32 {
103    unsafe {
104        use std::arch::aarch64::*;
105        debug_assert!(len.is_multiple_of(16));
106        let mut acc0 = vdupq_n_f32(0.0);
107        let mut acc1 = vdupq_n_f32(0.0);
108        let mut acc2 = vdupq_n_f32(0.0);
109        let mut acc3 = vdupq_n_f32(0.0);
110        let mut i = 0usize;
111        while i < len {
112            let a0 = vld1q_f32(a.add(i));
113            let a1 = vld1q_f32(a.add(i + 4));
114            let a2 = vld1q_f32(a.add(i + 8));
115            let a3 = vld1q_f32(a.add(i + 12));
116            let b0 = vld1q_f32(b.add(i));
117            let b1 = vld1q_f32(b.add(i + 4));
118            let b2 = vld1q_f32(b.add(i + 8));
119            let b3 = vld1q_f32(b.add(i + 12));
120            acc0 = vfmaq_f32(acc0, a0, b0);
121            acc1 = vfmaq_f32(acc1, a1, b1);
122            acc2 = vfmaq_f32(acc2, a2, b2);
123            acc3 = vfmaq_f32(acc3, a3, b3);
124            i += 16;
125        }
126        let acc = vaddq_f32(vaddq_f32(acc0, acc1), vaddq_f32(acc2, acc3));
127        vaddvq_f32(acc)
128    }
129}
130
131/// Same as [`dot_neon_multiple_of_16`] but with the length fixed at 128.
132/// Const bound lets the compiler unroll the loop fully — in benchmarks this
133/// specialisation shaves ~30% off the general-dimension path (~2.4 ms vs
134/// ~3.2 ms on the 250-candidate MaxSim rescoring kernel).
135#[cfg(target_arch = "aarch64")]
136#[inline(always)]
137unsafe fn dot_neon_128(a: *const f32, b: *const f32) -> f32 {
138    unsafe {
139        use std::arch::aarch64::*;
140        let mut acc0 = vdupq_n_f32(0.0);
141        let mut acc1 = vdupq_n_f32(0.0);
142        let mut acc2 = vdupq_n_f32(0.0);
143        let mut acc3 = vdupq_n_f32(0.0);
144        let mut i = 0usize;
145        while i < 128 {
146            let a0 = vld1q_f32(a.add(i));
147            let a1 = vld1q_f32(a.add(i + 4));
148            let a2 = vld1q_f32(a.add(i + 8));
149            let a3 = vld1q_f32(a.add(i + 12));
150            let b0 = vld1q_f32(b.add(i));
151            let b1 = vld1q_f32(b.add(i + 4));
152            let b2 = vld1q_f32(b.add(i + 8));
153            let b3 = vld1q_f32(b.add(i + 12));
154            acc0 = vfmaq_f32(acc0, a0, b0);
155            acc1 = vfmaq_f32(acc1, a1, b1);
156            acc2 = vfmaq_f32(acc2, a2, b2);
157            acc3 = vfmaq_f32(acc3, a3, b3);
158            i += 16;
159        }
160        let acc = vaddq_f32(vaddq_f32(acc0, acc1), vaddq_f32(acc2, acc3));
161        vaddvq_f32(acc)
162    }
163}
164
165/// ColBERT's late-interaction score: sum of per-query-token maxima.
166pub fn maxsim(query: &[Vector], document: &[Vector]) -> f32 {
167    if query.is_empty() || document.is_empty() {
168        return 0.0;
169    }
170    let document: Vec<_> = document.iter().map(|v| normalize(v)).collect();
171    query
172        .iter()
173        .map(|query| {
174            let query = normalize(query);
175            document
176                .iter()
177                .map(|doc| dot(&query, doc))
178                .fold(f32::NEG_INFINITY, f32::max)
179        })
180        .sum()
181}
182
183/// MaxSim over already-normalized query vectors and a flat document matrix.
184///
185/// Routes to a vectorised kernel when the dimension is a multiple of 16 on
186/// aarch64 (covers the standard ColBERT/E5/MPNet embedding sizes 128, 384,
187/// 512, 768, 1024), and to the packed AVX-512/AVX2 kernel for any dimension
188/// on x86_64. Falls back to the scalar path everywhere else.
189///
190/// When one query is scored against many documents, build a [`MaxSimQuery`]
191/// once instead: on x86_64 this call re-packs the query every time.
192pub fn maxsim_flat(query: &[Vector], document: &[f32], dimension: usize) -> f32 {
193    #[cfg(target_arch = "aarch64")]
194    {
195        if dimension > 0 && dimension.is_multiple_of(16) && dimension <= 4096 {
196            return maxsim_flat_neon(query, document, dimension);
197        }
198    }
199    #[cfg(target_arch = "x86_64")]
200    if dimension > 0 && bf16_active() {
201        return maxsim_flat_avx512_bf16(query, document, dimension);
202    }
203    #[cfg(target_arch = "x86_64")]
204    {
205        if let Some(kernel) = x86::PackedKernel::detect()
206            && x86::applicable(query, document, dimension)
207        {
208            thread_local! {
209                static PACKED: std::cell::RefCell<x86::Panel> =
210                    const { std::cell::RefCell::new(x86::Panel::new()) };
211            }
212            return PACKED.with(|cell| {
213                let mut packed = cell.borrow_mut();
214                x86::pack(query, dimension, kernel.lanes(), &mut packed);
215                // SAFETY: `detect` verified CPU support; `pack` sized the panel.
216                unsafe { kernel.score(&packed, query.len(), dimension, document) }
217            });
218        }
219    }
220    maxsim_flat_scalar(query, document, dimension)
221}
222
223/// A normalized query prepared for repeated MaxSim scoring.
224///
225/// On x86_64 with AVX2 or AVX-512 the query tokens are packed once into a
226/// dimension-major panel, so each document is scored with a register-tiled
227/// kernel that needs no horizontal reductions. Elsewhere it forwards to
228/// [`maxsim_flat`].
229///
230/// `score` always equals [`maxsim_flat`] on the same inputs. When the BF16 kernels are in use
231/// (see [`bf16_active`]) nothing is packed and every call goes through the same BF16 path, so
232/// the two never disagree about a document's score.
233pub struct MaxSimQuery<'a> {
234    tokens: &'a [Vector],
235    #[cfg_attr(not(target_arch = "x86_64"), allow(dead_code))]
236    dimension: usize,
237    #[cfg(target_arch = "x86_64")]
238    packed: Option<(x86::PackedKernel, x86::Panel)>,
239}
240
241impl<'a> MaxSimQuery<'a> {
242    pub fn new(tokens: &'a [Vector], dimension: usize) -> Self {
243        #[cfg(target_arch = "x86_64")]
244        let packed = x86::PackedKernel::detect()
245            .filter(|_| dimension > 0 && !tokens.is_empty() && !bf16_active())
246            .map(|kernel| {
247                let mut panel = x86::Panel::new();
248                x86::pack(tokens, dimension, kernel.lanes(), &mut panel);
249                (kernel, panel)
250            });
251        MaxSimQuery {
252            tokens,
253            dimension,
254            #[cfg(target_arch = "x86_64")]
255            packed,
256        }
257    }
258
259    /// Equivalent to `maxsim_flat(tokens, document, dimension)`.
260    #[inline]
261    pub fn score(&self, document: &[f32], dimension: usize) -> f32 {
262        #[cfg(target_arch = "x86_64")]
263        if let Some((kernel, panel)) = &self.packed
264            && dimension == self.dimension
265            && x86::applicable(self.tokens, document, dimension)
266        {
267            // SAFETY: `detect` verified CPU support; `pack` sized the panel.
268            return unsafe { kernel.score(panel, self.tokens.len(), dimension, document) };
269        }
270        maxsim_flat(self.tokens, document, dimension)
271    }
272}
273
274#[cfg(target_arch = "x86_64")]
275#[target_feature(enable = "avx512f,avx512bf16")]
276#[allow(unsafe_op_in_unsafe_fn)]
277unsafe fn dot_avx512_bf16_len(a: *const f32, b: *const f32, len: usize) -> f32 {
278    use std::arch::x86_64::*;
279    let mut acc0 = _mm512_setzero_ps();
280    let mut acc1 = _mm512_setzero_ps();
281    let mut acc2 = _mm512_setzero_ps();
282    let mut acc3 = _mm512_setzero_ps();
283    let mut i = 0usize;
284    while i + 128 <= len {
285        let a0 = _mm512_loadu_ps(a.add(i));
286        let a1 = _mm512_loadu_ps(a.add(i + 16));
287        let b0 = _mm512_loadu_ps(b.add(i));
288        let b1 = _mm512_loadu_ps(b.add(i + 16));
289        acc0 = _mm512_dpbf16_ps(
290            acc0,
291            _mm512_cvtne2ps_pbh(a1, a0),
292            _mm512_cvtne2ps_pbh(b1, b0),
293        );
294        let a2 = _mm512_loadu_ps(a.add(i + 32));
295        let a3 = _mm512_loadu_ps(a.add(i + 48));
296        let b2 = _mm512_loadu_ps(b.add(i + 32));
297        let b3 = _mm512_loadu_ps(b.add(i + 48));
298        acc1 = _mm512_dpbf16_ps(
299            acc1,
300            _mm512_cvtne2ps_pbh(a3, a2),
301            _mm512_cvtne2ps_pbh(b3, b2),
302        );
303        let a4 = _mm512_loadu_ps(a.add(i + 64));
304        let a5 = _mm512_loadu_ps(a.add(i + 80));
305        let b4 = _mm512_loadu_ps(b.add(i + 64));
306        let b5 = _mm512_loadu_ps(b.add(i + 80));
307        acc2 = _mm512_dpbf16_ps(
308            acc2,
309            _mm512_cvtne2ps_pbh(a5, a4),
310            _mm512_cvtne2ps_pbh(b5, b4),
311        );
312        let a6 = _mm512_loadu_ps(a.add(i + 96));
313        let a7 = _mm512_loadu_ps(a.add(i + 112));
314        let b6 = _mm512_loadu_ps(b.add(i + 96));
315        let b7 = _mm512_loadu_ps(b.add(i + 112));
316        acc3 = _mm512_dpbf16_ps(
317            acc3,
318            _mm512_cvtne2ps_pbh(a7, a6),
319            _mm512_cvtne2ps_pbh(b7, b6),
320        );
321        i += 128;
322    }
323    while i + 32 <= len {
324        let a0 = _mm512_loadu_ps(a.add(i));
325        let a1 = _mm512_loadu_ps(a.add(i + 16));
326        let b0 = _mm512_loadu_ps(b.add(i));
327        let b1 = _mm512_loadu_ps(b.add(i + 16));
328        acc0 = _mm512_dpbf16_ps(
329            acc0,
330            _mm512_cvtne2ps_pbh(a1, a0),
331            _mm512_cvtne2ps_pbh(b1, b0),
332        );
333        i += 32;
334    }
335    acc0 = _mm512_add_ps(acc0, acc1);
336    acc2 = _mm512_add_ps(acc2, acc3);
337    acc0 = _mm512_add_ps(acc0, acc2);
338    let mut result = _mm512_reduce_add_ps(acc0);
339    while i < len {
340        result += *a.add(i) * *b.add(i);
341        i += 1;
342    }
343    result
344}
345
346#[cfg(target_arch = "x86_64")]
347fn maxsim_flat_avx512_bf16(query: &[Vector], document: &[f32], dimension: usize) -> f32 {
348    let dot_fn: unsafe fn(*const f32, *const f32, usize) -> f32 = dot_avx512_bf16_len;
349    query
350        .iter()
351        .map(|q| {
352            debug_assert_eq!(q.len(), dimension);
353            let qp = q.as_ptr();
354            let mut best = f32::NEG_INFINITY;
355            for doc in document.chunks_exact(dimension) {
356                let s = unsafe { dot_fn(qp, doc.as_ptr(), dimension) };
357                if s > best {
358                    best = s;
359                }
360            }
361            best
362        })
363        .sum()
364}
365
366#[inline]
367fn maxsim_flat_scalar(query: &[Vector], document: &[f32], dimension: usize) -> f32 {
368    query
369        .iter()
370        .map(|q| {
371            document
372                .chunks_exact(dimension)
373                .map(|d| dot(q, d))
374                .fold(f32::NEG_INFINITY, f32::max)
375        })
376        .sum()
377}
378
379#[cfg(target_arch = "aarch64")]
380fn maxsim_flat_neon(query: &[Vector], document: &[f32], dimension: usize) -> f32 {
381    // NOTE: query-outer / doc-inner order is intentional. A previous attempt
382    // to swap the loops (stream doc tokens once, iterate 32 query tokens
383    // per doc) regressed the kernel from 2.35 ms to 3.51 ms in the
384    // maxsim-bench harness. The original order keeps the current query
385    // vector pinned in registers across all 200 doc-token dot products,
386    // and lets the compiler track a single scalar `best` in a register
387    // across the inner loop. Doc-outer loses both benefits.
388    if dimension == 128 {
389        // Const-length specialisation for the standard ColBERT dim so the
390        // compiler can fully unroll the inner FMA loop.
391        return query
392            .iter()
393            .map(|q| {
394                debug_assert_eq!(q.len(), 128);
395                let qp = q.as_ptr();
396                let mut best = f32::NEG_INFINITY;
397                for doc in document.as_chunks::<128>().0 {
398                    let s = unsafe { dot_neon_128(qp, doc.as_ptr()) };
399                    if s > best {
400                        best = s;
401                    }
402                }
403                best
404            })
405            .sum();
406    }
407    query
408        .iter()
409        .map(|q| {
410            debug_assert_eq!(q.len(), dimension);
411            let qp = q.as_ptr();
412            let mut best = f32::NEG_INFINITY;
413            for doc in document.chunks_exact(dimension) {
414                // SAFETY: `dimension` is validated a multiple of 16 by the caller
415                // (maxsim_flat), q and doc are contiguous slices of `dimension`
416                // f32 values, so the NEON loads stay in-bounds.
417                let s = unsafe { dot_neon_multiple_of_16(qp, doc.as_ptr(), dimension) };
418                if s > best {
419                    best = s;
420                }
421            }
422            best
423        })
424        .sum()
425}
426
427/// Packed MaxSim for x86_64.
428///
429/// MaxSim is a small GEMM (`query [nq x dim] * doc^T [dim x nd]`) followed by a
430/// row max. Instead of `nq * nd` independent dot products, each needing a
431/// horizontal reduction, the query is packed dimension-major into blocks of
432/// one vector register (16 lanes on AVX-512, 8 on AVX2), so lane `t` of a
433/// block holds query token `t`. A tile of `QB` query blocks x `DB` document
434/// tokens then accumulates in `QB * DB` registers: for each dimension `k` it
435/// loads `QB` query rows, broadcasts `DB` document scalars and issues
436/// `QB * DB` FMAs. The row max is a vertical `max` against a running best,
437/// so no shuffles are needed until the final per-block lane sum.
438#[cfg(target_arch = "x86_64")]
439mod x86 {
440    use super::Vector;
441    use annex::vector::kernels::Isa;
442
443    #[derive(Clone, Copy, Debug, PartialEq, Eq)]
444    pub(super) enum PackedKernel {
445        Avx2,
446        Avx512,
447    }
448
449    impl PackedKernel {
450        pub(super) fn detect() -> Option<Self> {
451            [PackedKernel::Avx512, PackedKernel::Avx2]
452                .into_iter()
453                .find(|k| k.is_supported())
454        }
455
456        pub(super) fn is_supported(self) -> bool {
457            match self {
458                PackedKernel::Avx2 => Isa::Avx2.is_supported(),
459                PackedKernel::Avx512 => std::arch::is_x86_feature_detected!("avx512f"),
460            }
461        }
462
463        pub(super) fn lanes(self) -> usize {
464            match self {
465                PackedKernel::Avx2 => 8,
466                PackedKernel::Avx512 => 16,
467            }
468        }
469
470        /// # Safety
471        /// The CPU must support `self`, and `panel` must come from [`pack`]
472        /// with the same `nq`, `dim` and `self.lanes()`.
473        pub(super) unsafe fn score(self, panel: &Panel, nq: usize, dim: usize, doc: &[f32]) -> f32 {
474            debug_assert!(panel.len() >= nq.div_ceil(self.lanes()) * dim * self.lanes());
475            let nd = doc.len() / dim;
476            unsafe {
477                match self {
478                    PackedKernel::Avx2 => avx2::maxsim(panel.as_ptr(), nq, dim, doc.as_ptr(), nd),
479                    PackedKernel::Avx512 => {
480                        avx512::maxsim(panel.as_ptr(), nq, dim, doc.as_ptr(), nd)
481                    }
482                }
483            }
484        }
485    }
486
487    /// The packed kernels need a query token and a whole document token;
488    /// everything else keeps the scalar path's semantics.
489    pub(super) fn applicable(query: &[Vector], document: &[f32], dimension: usize) -> bool {
490        dimension > 0 && !query.is_empty() && document.len() >= dimension
491    }
492
493    #[derive(Clone, Copy)]
494    #[repr(C, align(64))]
495    struct Line([f32; 16]);
496
497    /// Cache-line-aligned f32 buffer, so no packed query row splits a line.
498    #[derive(Default)]
499    pub(super) struct Panel(Vec<Line>);
500
501    impl Panel {
502        pub(super) const fn new() -> Self {
503            Panel(Vec::new())
504        }
505
506        fn len(&self) -> usize {
507            self.0.len() * 16
508        }
509
510        fn as_ptr(&self) -> *const f32 {
511            self.0.as_ptr().cast()
512        }
513
514        fn reset(&mut self, len: usize) -> &mut [f32] {
515            self.0.clear();
516            self.0.resize(len.div_ceil(16), Line([0.0; 16]));
517            // SAFETY: `Line` is 16 contiguous f32s with no padding.
518            unsafe { std::slice::from_raw_parts_mut(self.0.as_mut_ptr().cast(), self.len()) }
519        }
520    }
521
522    /// Pack `query` into blocks of `lanes` tokens, dimension-major: token
523    /// `t`, dimension `k` lands at `(t / lanes) * dim * lanes + k * lanes +
524    /// t % lanes`. Unused lanes and dimensions past a short token stay zero,
525    /// matching `dot`'s shorter-length semantics.
526    pub(super) fn pack(query: &[Vector], dim: usize, lanes: usize, panel: &mut Panel) {
527        let out = panel.reset(query.len().div_ceil(lanes) * dim * lanes);
528        for (block, tokens) in out.chunks_exact_mut(dim * lanes).zip(query.chunks(lanes)) {
529            if tokens.len() == lanes && tokens.iter().all(|t| t.len() >= dim) {
530                // Full block: write each output row sequentially.
531                for (k, row) in block.chunks_exact_mut(lanes).enumerate() {
532                    for (slot, token) in row.iter_mut().zip(tokens) {
533                        // SAFETY: every token has at least `dim` values.
534                        *slot = unsafe { *token.get_unchecked(k) };
535                    }
536                }
537            } else {
538                for (lane, token) in tokens.iter().enumerate() {
539                    for (k, &v) in token.iter().take(dim).enumerate() {
540                        block[k * lanes + lane] = v;
541                    }
542                }
543            }
544        }
545    }
546
547    macro_rules! packed_maxsim {
548        (
549            $name:ident, $feature:literal, $reg:ty, $lanes:literal,
550            $db_big:literal, $db_mid:literal, $qb_max:literal,
551            zero: $zero:expr, splat: $splat:path, load: $load:path,
552            fma: $fma:path, max: $max:path, sum: $sum:ident
553        ) => {
554            mod $name {
555                use std::arch::x86_64::*;
556                const L: usize = $lanes;
557
558                #[target_feature(enable = $feature)]
559                pub(super) unsafe fn maxsim(
560                    panel: *const f32,
561                    nq: usize,
562                    dim: usize,
563                    doc: *const f32,
564                    nd: usize,
565                ) -> f32 {
566                    let blocks = nq.div_ceil(L);
567                    let mut total = 0.0f32;
568                    let mut b = 0;
569                    while b < blocks {
570                        let q = unsafe { panel.add(b * dim * L) };
571                        if b + $qb_max <= blocks {
572                            let best = unsafe { sweep::<$qb_max>(q, dim, doc, nd) };
573                            for (i, v) in best.into_iter().enumerate() {
574                                total += unsafe { $sum(v, nq - (b + i) * L) };
575                            }
576                            b += $qb_max;
577                        } else {
578                            let best = unsafe { sweep::<1>(q, dim, doc, nd) };
579                            total += unsafe { $sum(best[0], nq - b * L) };
580                            b += 1;
581                        }
582                    }
583                    total
584                }
585
586                #[inline]
587                #[target_feature(enable = $feature)]
588                unsafe fn sweep<const QB: usize>(
589                    q: *const f32,
590                    dim: usize,
591                    doc: *const f32,
592                    nd: usize,
593                ) -> [$reg; QB] {
594                    let mut best = [$splat(f32::NEG_INFINITY); QB];
595                    let mut d = 0;
596                    unsafe {
597                        while d + $db_big <= nd {
598                            tile::<QB, $db_big>(q, dim, doc.add(d * dim), &mut best);
599                            d += $db_big;
600                        }
601                        if d + $db_mid <= nd {
602                            tile::<QB, $db_mid>(q, dim, doc.add(d * dim), &mut best);
603                            d += $db_mid;
604                        }
605                        while d < nd {
606                            tile::<QB, 1>(q, dim, doc.add(d * dim), &mut best);
607                            d += 1;
608                        }
609                    }
610                    best
611                }
612
613                /// `QB` query blocks x `DB` document tokens, all in registers.
614                #[inline]
615                #[target_feature(enable = $feature)]
616                unsafe fn tile<const QB: usize, const DB: usize>(
617                    q: *const f32,
618                    dim: usize,
619                    docs: *const f32,
620                    best: &mut [$reg; QB],
621                ) {
622                    let mut acc = [[$zero; QB]; DB];
623                    for k in 0..dim {
624                        let mut qv = [$zero; QB];
625                        for (b, v) in qv.iter_mut().enumerate() {
626                            *v = unsafe { $load(q.add(b * dim * L + k * L)) };
627                        }
628                        for (j, row) in acc.iter_mut().enumerate() {
629                            let x = $splat(unsafe { *docs.add(j * dim + k) });
630                            for (a, &v) in row.iter_mut().zip(qv.iter()) {
631                                *a = $fma(v, x, *a);
632                            }
633                        }
634                    }
635                    for row in &acc {
636                        for (m, &a) in best.iter_mut().zip(row) {
637                            *m = $max(*m, a);
638                        }
639                    }
640                }
641
642                #[allow(dead_code)]
643                #[inline]
644                #[target_feature(enable = "avx512f")]
645                unsafe fn sum512(v: __m512, valid: usize) -> f32 {
646                    let m = if valid >= 16 {
647                        u16::MAX
648                    } else {
649                        ((1u32 << valid) - 1) as u16
650                    };
651                    _mm512_mask_reduce_add_ps(m, v)
652                }
653
654                #[allow(dead_code)]
655                #[inline]
656                #[target_feature(enable = "avx")]
657                unsafe fn sum256(v: __m256, valid: usize) -> f32 {
658                    let mut lanes = [0.0f32; 8];
659                    unsafe { _mm256_storeu_ps(lanes.as_mut_ptr(), v) };
660                    lanes[..valid.min(8)].iter().sum()
661                }
662            }
663        };
664    }
665
666    // AVX-512: 2 query blocks (32 tokens) x 8 document tokens = 16 zmm
667    // accumulators + 2 query rows + 1 broadcast, of 32 registers.
668    packed_maxsim!(
669        avx512, "avx512f", __m512, 16, 8, 4, 2,
670        zero: _mm512_setzero_ps(), splat: _mm512_set1_ps, load: _mm512_load_ps,
671        fma: _mm512_fmadd_ps, max: _mm512_max_ps, sum: sum512
672    );
673
674    // AVX2: 2 query blocks (16 tokens) x 6 document tokens = 12 ymm
675    // accumulators + 2 query rows + 1 broadcast, of 16 registers.
676    packed_maxsim!(
677        avx2, "avx2,fma", __m256, 8, 6, 3, 2,
678        zero: _mm256_setzero_ps(), splat: _mm256_set1_ps, load: _mm256_load_ps,
679        fma: _mm256_fmadd_ps, max: _mm256_max_ps, sum: sum256
680    );
681}
682
683#[cfg(test)]
684mod tests {
685    use super::*;
686
687    /// Round an f32 to the nearest bfloat16 (ties to even) and back, as `vcvtneps2bf16` does to
688    /// each operand before the fused multiply-accumulate.
689    fn bf16_round(x: f32) -> f32 {
690        let bits = x.to_bits();
691        f32::from_bits(bits.wrapping_add(0x7FFF + ((bits >> 16) & 1)) & 0xFFFF_0000)
692    }
693
694    fn deterministic(seed: u64, dim: usize) -> Vector {
695        let mut s = seed.wrapping_mul(0x9E37_79B9_7F4A_7C15);
696        (0..dim)
697            .map(|_| {
698                s = s
699                    .wrapping_mul(6364136223846793005)
700                    .wrapping_add(1442695040888963407);
701                (((s >> 33) as u32) as f32 / u32::MAX as f32) * 2.0 - 1.0
702            })
703            .collect()
704    }
705
706    fn flat(doc: &[Vector]) -> (Vec<f32>, usize) {
707        let dim = doc[0].len();
708        (doc.iter().flat_map(|v| v.iter().copied()).collect(), dim)
709    }
710
711    #[test]
712    fn maxsim_flat_scalar_and_dispatch_agree_on_dim128() {
713        let query: Vec<_> = (0..8).map(|i| deterministic(0x11 + i, 128)).collect();
714        let doc_tokens: Vec<_> = (0..50).map(|i| deterministic(0x2000 + i, 128)).collect();
715        let (flat_doc, dim) = flat(&doc_tokens);
716        let s = maxsim_flat_scalar(&query, &flat_doc, dim);
717        let d = maxsim_flat(&query, &flat_doc, dim);
718        assert!((s - d).abs() < 1e-3, "scalar={s} dispatch={d}");
719    }
720
721    #[test]
722    fn maxsim_flat_scalar_and_dispatch_agree_on_dim384() {
723        let query: Vec<_> = (0..12).map(|i| deterministic(0x33 + i, 384)).collect();
724        let doc_tokens: Vec<_> = (0..30).map(|i| deterministic(0x4000 + i, 384)).collect();
725        let (flat_doc, dim) = flat(&doc_tokens);
726        let s = maxsim_flat_scalar(&query, &flat_doc, dim);
727        let d = maxsim_flat(&query, &flat_doc, dim);
728        assert!((s - d).abs() < 1e-3, "scalar={s} dispatch={d}");
729    }
730
731    #[test]
732    fn maxsim_flat_dispatch_falls_back_when_dim_not_multiple_of_16() {
733        // 100 is not a multiple of 16 — dispatch must take the scalar path.
734        let query: Vec<_> = (0..4).map(|i| deterministic(0x55 + i, 100)).collect();
735        let doc_tokens: Vec<_> = (0..10).map(|i| deterministic(0x6000 + i, 100)).collect();
736        let (flat_doc, dim) = flat(&doc_tokens);
737        let s = maxsim_flat_scalar(&query, &flat_doc, dim);
738        let d = maxsim_flat(&query, &flat_doc, dim);
739        // Exact on aarch64 (scalar path); x86_64 vectorises every dimension.
740        let tol = if cfg!(target_arch = "aarch64") {
741            1e-6
742        } else {
743            1e-4
744        };
745        assert!((s - d).abs() < tol, "scalar={s} dispatch={d}");
746    }
747
748    #[test]
749    fn maxsim_kernels_match_scalar_across_shapes() {
750        for &dim in &[1usize, 7, 16, 33, 96, 100, 128, 129, 384, 768] {
751            for &nq in &[1usize, 3, 8, 15, 16, 17, 31, 32, 33, 48] {
752                for &nd in &[1usize, 2, 5, 6, 7, 8, 9, 13, 64, 201] {
753                    let query: Vec<_> = (0..nq)
754                        .map(|i| normalize(&deterministic(0x77 + i as u64, dim)))
755                        .collect();
756                    let doc: Vec<_> = (0..nd)
757                        .map(|i| normalize(&deterministic(0x9000 + (i * 31 + dim) as u64, dim)))
758                        .collect();
759                    let (flat_doc, _) = flat(&doc);
760                    let want = maxsim_flat_scalar(&query, &flat_doc, dim);
761                    let tol = 1e-5 * nq as f32 + 1e-4;
762                    let got = maxsim_flat(&query, &flat_doc, dim);
763                    assert!(
764                        (got - want).abs() <= tol,
765                        "flat dim={dim} nq={nq} nd={nd}: {got} vs {want}"
766                    );
767                    let prepared = MaxSimQuery::new(&query, dim).score(&flat_doc, dim);
768                    assert!(
769                        (prepared - want).abs() <= tol,
770                        "prepared dim={dim} nq={nq} nd={nd}: {prepared} vs {want}"
771                    );
772                }
773            }
774        }
775    }
776
777    #[cfg(target_arch = "x86_64")]
778    #[test]
779    fn every_x86_maxsim_kernel_matches_scalar() {
780        for kernel in [x86::PackedKernel::Avx2, x86::PackedKernel::Avx512] {
781            if !kernel.is_supported() {
782                continue;
783            }
784            for &(dim, nq, nd) in &[(128, 32, 200), (100, 17, 13), (384, 5, 7), (1, 1, 1)] {
785                let query: Vec<_> = (0..nq)
786                    .map(|i| normalize(&deterministic(0x5 + i as u64, dim)))
787                    .collect();
788                let doc: Vec<_> = (0..nd)
789                    .map(|i| normalize(&deterministic(0x700 + i as u64, dim)))
790                    .collect();
791                let (flat_doc, _) = flat(&doc);
792                let mut panel = x86::Panel::new();
793                x86::pack(&query, dim, kernel.lanes(), &mut panel);
794                let got = unsafe { kernel.score(&panel, nq, dim, &flat_doc) };
795                let want = maxsim_flat_scalar(&query, &flat_doc, dim);
796                assert!(
797                    (got - want).abs() < 1e-3,
798                    "{kernel:?} dim={dim}: {got} vs {want}"
799                );
800            }
801        }
802    }
803
804    #[test]
805    fn maxsim_edge_cases_match_scalar() {
806        let q = vec![normalize(&deterministic(1, 16))];
807        assert_eq!(maxsim_flat(&[], &[1.0; 16], 16), 0.0);
808        assert_eq!(maxsim_flat(&q, &[], 16), f32::NEG_INFINITY);
809        assert_eq!(MaxSimQuery::new(&q, 16).score(&[], 16), f32::NEG_INFINITY);
810        // A partial trailing token is ignored, as with `chunks_exact`.
811        let doc = deterministic(2, 20);
812        let want = maxsim_flat_scalar(&q, &doc, 16);
813        assert!((maxsim_flat(&q, &doc, 16) - want).abs() < 1e-5);
814    }
815
816    #[test]
817    fn maxsim_flat_bf16_agrees_with_exact_f64_on_dim128() {
818        // Only meaningful when the BF16 kernels are in use (CPU support plus VECTORDB_BF16=1).
819        if !bf16_active() {
820            return;
821        }
822        // Compare with a plain f64 reference, not `maxsim_flat_scalar`: that calls `dot`, which
823        // is the BF16 kernel here, so it would only compare BF16 with itself.
824        let mut rng = 0xcafe_babe_u64;
825        let mut next = || -> f32 {
826            rng ^= rng << 13;
827            rng ^= rng >> 7;
828            rng ^= rng << 17;
829            (rng as f32 / u64::MAX as f32) * 2.0 - 1.0
830        };
831        let query: Vec<_> = (0..8)
832            .map(|_| normalize(&(0..128).map(|_| next()).collect::<Vec<f32>>()))
833            .collect();
834        let doc_tokens: Vec<_> = (0..50)
835            .map(|_| normalize(&(0..128).map(|_| next()).collect::<Vec<f32>>()))
836            .collect();
837        let flat_doc: Vec<f32> = doc_tokens.iter().flat_map(|v| v.iter().copied()).collect();
838        let exact: f64 = query
839            .iter()
840            .map(|q| {
841                doc_tokens
842                    .iter()
843                    .map(|d| {
844                        q.iter()
845                            .zip(d)
846                            .map(|(&x, &y)| f64::from(x) * f64::from(y))
847                            .sum::<f64>()
848                    })
849                    .fold(f64::NEG_INFINITY, f64::max)
850            })
851            .sum();
852        let dispatch = maxsim_flat(&query, &flat_doc, 128);
853        // One BF16 dot error per query token at most.
854        let tol = query.len() as f64 * f64::from(BF16_UNIT_DOT_TOL);
855        assert!(
856            (f64::from(dispatch) - exact).abs() <= tol,
857            "bf16 maxsim {dispatch} vs exact {exact} (tolerance {tol})"
858        );
859    }
860
861    /// `MaxSimQuery::score` is documented as equal to `maxsim_flat`. That has to hold on every
862    /// CPU and with BF16 on or off, because retrieval scores through the former and exact
863    /// scoring through the latter.
864    #[test]
865    fn prepared_query_scores_exactly_like_maxsim_flat() {
866        for &dim in &[1usize, 16, 33, 100, 128, 384] {
867            for &nq in &[1usize, 4, 17] {
868                let query: Vec<_> = (0..nq)
869                    .map(|i| normalize(&deterministic(0x31 + i as u64, dim)))
870                    .collect();
871                let doc: Vec<_> = (0..9)
872                    .map(|i| normalize(&deterministic(0x7000 + i as u64, dim)))
873                    .collect();
874                let (flat_doc, _) = flat(&doc);
875                let prepared = MaxSimQuery::new(&query, dim).score(&flat_doc, dim);
876                let direct = maxsim_flat(&query, &flat_doc, dim);
877                assert_eq!(
878                    prepared.to_bits(),
879                    direct.to_bits(),
880                    "dim={dim} nq={nq}: prepared {prepared} != maxsim_flat {direct}"
881                );
882            }
883        }
884    }
885
886    #[test]
887    fn dot_self_is_near_one_after_normalize_on_bf16() {
888        // Only meaningful when the BF16 kernels are in use (CPU support plus VECTORDB_BF16=1).
889        if !bf16_active() {
890            return;
891        }
892        for dim in [64usize, 128, 384, 768] {
893            let raw: Vec<f32> = (0..dim).map(|i| (i as f32 + 1.0).recip()).collect();
894            let normed = normalize(&raw);
895            let self_dot = dot(&normed, &normed);
896            assert!(
897                (self_dot - 1.0).abs() < BF16_UNIT_DOT_TOL,
898                "dim={dim}: self_dot={self_dot}"
899            );
900        }
901    }
902
903    /// Runs on every CPU. Emulates BF16 operand rounding in software to check the documented
904    /// error bound, so the BF16 tolerances above are validated without BF16 hardware.
905    #[test]
906    fn bf16_operand_rounding_stays_within_the_documented_bound() {
907        let emulated_dot = |a: &[f32], b: &[f32]| -> f64 {
908            a.iter()
909                .zip(b)
910                .map(|(&x, &y)| f64::from(bf16_round(x)) * f64::from(bf16_round(y)))
911                .sum()
912        };
913        let exact_dot = |a: &[f32], b: &[f32]| -> f64 {
914            a.iter()
915                .zip(b)
916                .map(|(&x, &y)| f64::from(x) * f64::from(y))
917                .sum()
918        };
919
920        // Reference case: the self-dot of a normalised 1/(i+1) vector at dim 64 was observed as
921        // 1.0038899 on a BF16 runner, which is why 1e-3 cannot be the tolerance.
922        let raw: Vec<f32> = (0..64).map(|i| (i as f32 + 1.0).recip()).collect();
923        let normed = normalize(&raw);
924        let self_dot = emulated_dot(&normed, &normed);
925        assert!(
926            (self_dot - 1.0038899).abs() < 1e-5,
927            "emulation no longer matches the hardware observation: {self_dot}"
928        );
929        assert!((self_dot - 1.0).abs() > 1e-3);
930
931        let mut worst = 0.0f64;
932        for &dim in &[1usize, 7, 16, 33, 64, 100, 128, 129, 384, 768] {
933            let vectors: Vec<_> = (0..24)
934                .map(|i| normalize(&deterministic(0xB0 + i, dim)))
935                .collect();
936            for a in &vectors {
937                for b in &vectors {
938                    let error = (emulated_dot(a, b) - exact_dot(a, b)).abs();
939                    worst = worst.max(error);
940                    assert!(
941                        error < f64::from(BF16_UNIT_DOT_TOL),
942                        "dim={dim}: BF16 rounding error {error} exceeds the bound"
943                    );
944                }
945            }
946        }
947        // The bound is tight enough to mean something: real errors reach a good fraction of it.
948        assert!(
949            worst > 1e-3,
950            "worst observed error {worst} is suspiciously small"
951        );
952    }
953}