1#[cfg(target_arch = "x86_64")]
2use annex::vector::simd::{CpuLevel, cpu_level};
3
4pub type Vector = Vec<f32>;
5
6pub(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#[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#[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#[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 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 annex::vector::kernels::dot(&left[..n], &right[..n])
95 }
96}
97
98#[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#[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
165pub 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
183pub 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 unsafe { kernel.score(&packed, query.len(), dimension, document) }
217 });
218 }
219 }
220 maxsim_flat_scalar(query, document, dimension)
221}
222
223pub 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 #[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 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 if dimension == 128 {
389 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 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#[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 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 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 #[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 unsafe { std::slice::from_raw_parts_mut(self.0.as_mut_ptr().cast(), self.len()) }
519 }
520 }
521
522 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 for (k, row) in block.chunks_exact_mut(lanes).enumerate() {
532 for (slot, token) in row.iter_mut().zip(tokens) {
533 *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 #[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 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 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 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 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 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 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 if !bf16_active() {
820 return;
821 }
822 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 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 #[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 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 #[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 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 assert!(
949 worst > 1e-3,
950 "worst observed error {worst} is suspiciously small"
951 );
952 }
953}