1use std::hint::black_box;
32use std::sync::OnceLock;
33use std::time::Instant;
34
35use crate::lut::Lut;
36use crate::query::PreparedQuery;
37use crate::scorer::Codes;
38use crate::MAX_DIM;
39
40#[cfg(target_arch = "x86_64")]
41mod avx2;
42#[cfg(target_arch = "x86_64")]
43mod avx2_vnni;
44#[cfg(target_arch = "x86_64")]
45mod avx512;
46#[cfg(target_arch = "aarch64")]
47mod neon;
48#[cfg(target_arch = "aarch64")]
49mod neon_i8mm;
50
51#[derive(Debug, Clone, Copy, PartialEq, Eq)]
54pub enum Kernel {
55 Scalar(ScalarReason),
57 NeonSdot,
59 NeonI8mm,
63 Avx2,
65 Avx2Vnni,
68 Avx512Vnni,
70}
71
72#[derive(Debug, Clone, Copy, PartialEq, Eq)]
74pub enum ScalarReason {
75 Forced,
77 DimNotSimdAligned,
79 NoNibbleTables,
81 CpuUnsupported,
84 SelfCheckFailed,
89}
90
91impl std::fmt::Display for Kernel {
92 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
93 match self {
94 Kernel::Scalar(r) => write!(f, "scalar ({r:?})"),
95 Kernel::NeonSdot => write!(f, "neon-sdot"),
96 Kernel::NeonI8mm => write!(f, "neon-i8mm"),
97 Kernel::Avx2 => write!(f, "avx2"),
98 Kernel::Avx2Vnni => write!(f, "avx2-vnni"),
99 Kernel::Avx512Vnni => write!(f, "avx512-vnni"),
100 }
101 }
102}
103
104impl Kernel {
105 pub fn is_simd(&self) -> bool {
107 !matches!(self, Kernel::Scalar(_))
108 }
109}
110
111fn env_force_scalar() -> bool {
112 static FORCE: OnceLock<bool> = OnceLock::new();
113 *FORCE.get_or_init(|| {
114 std::env::var("MAXSIM_LUT_FORCE_SCALAR")
115 .map(|v| v == "1" || v.eq_ignore_ascii_case("true"))
116 .unwrap_or(false)
117 })
118}
119
120fn env_pin() -> Option<Kernel> {
126 static PIN: OnceLock<Option<Kernel>> = OnceLock::new();
127 *PIN.get_or_init(|| match std::env::var("MAXSIM_LUT_KERNEL").ok()?.as_str() {
128 "neon-sdot" => Some(Kernel::NeonSdot),
129 "neon-i8mm" => Some(Kernel::NeonI8mm),
130 "avx2" => Some(Kernel::Avx2),
131 "avx2-vnni" => Some(Kernel::Avx2Vnni),
132 "avx512-vnni" => Some(Kernel::Avx512Vnni),
133 _ => None,
134 })
135}
136
137fn cpu_kernels() -> &'static [Kernel] {
141 #[cfg(target_arch = "aarch64")]
142 {
143 let dotprod = std::arch::is_aarch64_feature_detected!("dotprod");
144 let i8mm = std::arch::is_aarch64_feature_detected!("i8mm");
145 match (dotprod, i8mm) {
146 (true, true) => &[Kernel::NeonI8mm, Kernel::NeonSdot],
147 (true, false) => &[Kernel::NeonSdot],
148 (false, true) => &[Kernel::NeonI8mm],
149 (false, false) => &[],
150 }
151 }
152 #[cfg(target_arch = "x86_64")]
153 {
154 let avx512 = is_x86_feature_detected!("avx512f")
155 && is_x86_feature_detected!("avx512bw")
156 && is_x86_feature_detected!("avx512vnni");
157 let vnni256 = is_x86_feature_detected!("avxvnni");
158 let avx2 = is_x86_feature_detected!("avx2");
159 match (avx512, vnni256, avx2) {
160 (true, true, _) => &[Kernel::Avx512Vnni, Kernel::Avx2Vnni, Kernel::Avx2],
161 (true, false, _) => &[Kernel::Avx512Vnni, Kernel::Avx2],
162 (false, true, _) => &[Kernel::Avx2Vnni, Kernel::Avx2],
163 (false, false, true) => &[Kernel::Avx2],
164 (false, false, false) => &[],
165 }
166 }
167 #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
168 &[]
169}
170
171fn env_no_calibrate() -> bool {
175 static SKIP: OnceLock<bool> = OnceLock::new();
176 *SKIP.get_or_init(|| {
177 std::env::var("MAXSIM_LUT_NO_CALIBRATE")
178 .map(|v| v == "1" || v.eq_ignore_ascii_case("true"))
179 .unwrap_or(false)
180 })
181}
182
183struct Probe {
185 lut: Lut,
186 query: PreparedQuery,
187 packed: Vec<u8>,
188 inv: Vec<f32>,
189 row_stride: usize,
190 n_tokens: usize,
191}
192
193impl Probe {
194 fn new(nbits: usize, dim: usize, nq: usize, n_tokens: usize) -> Option<Self> {
197 let n = 1usize << nbits;
198 let weights: Vec<f32> = (0..n)
199 .map(|i| -0.35 + 0.7 * (i as f32 + 0.5) / n as f32)
200 .collect();
201 let lut = Lut::colbert(nbits, &weights).ok()?;
202 let q: Vec<f32> = (0..nq * dim)
206 .map(|i| ((i * 37 % 251) as f32 / 251.0) - 0.5)
207 .collect();
208 let query = PreparedQuery::new(&lut, &q, nq, dim).ok()?;
209 let row_stride = dim / lut.keys_per_byte();
210 Some(Self {
211 packed: (0..n_tokens * row_stride).map(|i| (i * 97 % 256) as u8).collect(),
212 inv: (0..n_tokens).map(|i| 0.8 + (i % 7) as f32 * 0.05).collect(),
213 lut,
214 query,
215 row_stride,
216 n_tokens,
217 })
218 }
219
220 fn args(&self) -> Args<'_> {
221 Args {
222 query: &self.query,
223 packed: &self.packed,
224 row_stride: self.row_stride,
225 n_tokens: self.n_tokens,
226 codes: Codes::None,
227 cdot: self.query.zeros(),
228 cdot_stride: 0,
229 inv_norms: Some(&self.inv),
230 }
231 }
232}
233
234const PROBE_SHAPES: [(usize, usize, usize, usize); 4] = [
240 (4, 128, 32, 64),
242 (4, 40, 9, 5),
243 (2, 96, 17, 7),
244 (1, 256, 1, 3),
245];
246
247fn calibrated() -> Kernel {
260 static CHOICE: OnceLock<Kernel> = OnceLock::new();
261 *CHOICE.get_or_init(|| {
262 let cands = cpu_kernels();
263 match cands.first() {
264 None => Kernel::Scalar(ScalarReason::CpuUnsupported),
265 Some(&first) if env_no_calibrate() => first,
266 Some(&first) => measure_fastest(cands, first),
267 }
268 })
269}
270
271fn verified_kernels<F>(cands: &[Kernel], probes: &[Probe], mut score: F) -> Vec<Kernel>
275where
276 F: FnMut(Kernel, &Lut, &Args<'_>) -> f32,
277{
278 cands
279 .iter()
280 .copied()
281 .filter(|&k| {
282 probes.iter().all(|p| {
283 let args = p.args();
284 score(k, &p.lut, &args).to_bits() == scalar(&p.lut, &args).to_bits()
285 })
286 })
287 .collect()
288}
289
290fn measure_fastest(cands: &[Kernel], fallback: Kernel) -> Kernel {
297 const REPS: usize = 5;
298 const ITERS: usize = 8;
299
300 let mut probes = Vec::with_capacity(PROBE_SHAPES.len());
301 for (nbits, dim, nq, ntok) in PROBE_SHAPES {
302 match Probe::new(nbits, dim, nq, ntok) {
303 Some(p) => probes.push(p),
304 None => return fallback,
305 }
306 }
307
308 let verified = verified_kernels(cands, &probes, run);
309 debug_assert_eq!(
312 verified.len(),
313 cands.len(),
314 "a supported kernel disagreed with the scalar reference: kept {verified:?} of {cands:?}"
315 );
316 let Some((&first, rest)) = verified.split_first() else {
317 return Kernel::Scalar(ScalarReason::SelfCheckFailed);
318 };
319 if rest.is_empty() {
320 return first;
321 }
322
323 let timed = &probes[0];
324 let args = timed.args();
325 let mut best = vec![f64::INFINITY; verified.len()];
326 for _ in 0..REPS {
327 for (slot, &k) in best.iter_mut().zip(&verified) {
328 let t = Instant::now();
329 for _ in 0..ITERS {
330 black_box(run(k, &timed.lut, &args));
331 }
332 *slot = slot.min(t.elapsed().as_secs_f64());
333 }
334 }
335 let mut winner = 0usize;
336 for i in 1..verified.len() {
337 if best[i] < best[winner] {
338 winner = i;
339 }
340 }
341 verified[winner]
342}
343
344pub(crate) fn select(lut: &Lut, dim: usize) -> Kernel {
347 if lut.force_scalar_set() || env_force_scalar() {
348 return Kernel::Scalar(ScalarReason::Forced);
349 }
350 if !dim.is_multiple_of(8) || dim > MAX_DIM {
351 return Kernel::Scalar(ScalarReason::DimNotSimdAligned);
352 }
353 if lut.nibble_tables().is_none() {
354 return Kernel::Scalar(ScalarReason::NoNibbleTables);
355 }
356 for pin in [lut.pinned_kernel(), env_pin()].into_iter().flatten() {
359 if cpu_kernels().contains(&pin) {
360 return pin;
361 }
362 }
363 calibrated()
364}
365
366pub fn supported_kernels() -> &'static [Kernel] {
372 cpu_kernels()
373}
374
375pub fn warm_up() -> Kernel {
390 calibrated()
391}
392
393pub(crate) struct Args<'a> {
395 pub query: &'a PreparedQuery,
396 pub packed: &'a [u8],
398 pub row_stride: usize,
399 pub n_tokens: usize,
400 pub codes: Codes<'a>,
402 pub cdot: &'a [f32],
404 pub cdot_stride: usize,
405 pub inv_norms: Option<&'a [f32]>,
406}
407
408impl Args<'_> {
409 #[inline(always)]
410 pub(crate) fn inv(&self, t: usize) -> f32 {
411 match self.inv_norms {
412 Some(v) => v[t],
413 None => 1.0,
414 }
415 }
416 #[inline(always)]
417 pub(crate) fn crow(&self, t: usize) -> *const f32 {
418 let off = self.codes.id(t) * self.cdot_stride;
420 debug_assert!(off + self.query.n_tokens() <= self.cdot.len());
421 unsafe { self.cdot.as_ptr().add(off) }
423 }
424}
425
426#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
430thread_local! {
431 static SCRATCH: std::cell::RefCell<(Vec<f32>, Vec<i32>)> =
432 const { std::cell::RefCell::new((Vec::new(), Vec::new())) };
433}
434
435pub(crate) fn maxsim(lut: &Lut, args: &Args<'_>) -> f32 {
437 run(select(lut, args.query.dim()), lut, args)
438}
439
440#[cfg(test)]
445fn kernels_under_test() -> Vec<Kernel> {
446 let mut v = vec![Kernel::Scalar(ScalarReason::Forced)];
447 v.extend_from_slice(cpu_kernels());
448 v
449}
450
451pub(crate) fn run(kernel: Kernel, lut: &Lut, args: &Args<'_>) -> f32 {
455 match kernel {
456 Kernel::Scalar(_) => scalar(lut, args),
457 #[cfg(target_arch = "aarch64")]
458 Kernel::NeonSdot => SCRATCH.with(|s| {
459 let (best, accs) = &mut *s.borrow_mut();
460 unsafe { neon::maxsim(lut, args, best, accs) }
463 }),
464 #[cfg(target_arch = "aarch64")]
465 Kernel::NeonI8mm => SCRATCH.with(|s| {
466 let (best, accs) = &mut *s.borrow_mut();
467 unsafe { neon_i8mm::maxsim(lut, args, best, accs) }
469 }),
470 #[cfg(target_arch = "x86_64")]
471 Kernel::Avx2 => SCRATCH.with(|s| {
472 let (best, accs) = &mut *s.borrow_mut();
473 unsafe { avx2::maxsim(lut, args, best, accs) }
475 }),
476 #[cfg(target_arch = "x86_64")]
477 Kernel::Avx2Vnni => SCRATCH.with(|s| {
478 let (best, accs) = &mut *s.borrow_mut();
479 unsafe { avx2_vnni::maxsim(lut, args, best, accs) }
481 }),
482 #[cfg(target_arch = "x86_64")]
483 Kernel::Avx512Vnni => SCRATCH.with(|s| {
484 let (best, accs) = &mut *s.borrow_mut();
485 unsafe { avx512::maxsim(lut, args, best, accs) }
487 }),
488 #[allow(unreachable_patterns)]
489 _ => scalar(lut, args),
490 }
491}
492
493pub(crate) fn scalar(lut: &Lut, a: &Args<'_>) -> f32 {
496 let q = a.query;
497 let nq = q.n_tokens();
498 let dim = q.dim();
499 if nq == 0 || a.n_tokens == 0 {
500 return 0.0;
501 }
502 let kpb = lut.keys_per_byte();
503 let pdim = dim / kpb;
504 let qv = q.codes();
505 let sqw = q.sqw();
506 let mut best = vec![f32::NEG_INFINITY; nq];
507 let mut w = [0i8; MAX_DIM];
508 for t in 0..a.n_tokens {
509 let row = &a.packed[t * a.row_stride..t * a.row_stride + pdim];
510 for (i, &byte) in row.iter().enumerate() {
511 w[i * kpb..(i + 1) * kpb].copy_from_slice(lut.expand(byte));
512 }
513 let inv = a.inv(t);
514 let crow = a.crow(t);
515 for (qi, best_q) in best.iter_mut().enumerate() {
516 let qrow = &qv[qi * dim..(qi + 1) * dim];
517 let mut acc = 0i32;
518 for (qd, wd) in qrow.iter().zip(&w[..dim]) {
519 acc += *qd as i32 * *wd as i32;
520 }
521 let c = unsafe { *crow.add(qi) };
523 let score = (sqw[qi] * acc as f32 + c) * inv;
524 if score > *best_q {
525 *best_q = score;
526 }
527 }
528 }
529 best.iter().sum()
530}
531
532#[cfg(test)]
533mod tests {
534 use super::*;
535 use crate::ColbertPacking;
536
537 struct Rng(u64);
538 impl Rng {
539 fn next(&mut self) -> u64 {
540 let mut x = self.0;
541 x ^= x << 13;
542 x ^= x >> 7;
543 x ^= x << 17;
544 self.0 = x;
545 x
546 }
547 fn f32(&mut self, lo: f32, hi: f32) -> f32 {
548 lo + (hi - lo) * ((self.next() >> 40) as f32 / (1u64 << 24) as f32)
549 }
550 }
551
552 #[test]
558 fn the_self_check_rejects_a_kernel_that_disagrees() {
559 let cands = cpu_kernels();
560 if cands.is_empty() {
561 return; }
563 let probes: Vec<Probe> = PROBE_SHAPES
564 .iter()
565 .map(|&(nbits, dim, nq, ntok)| Probe::new(nbits, dim, nq, ntok).expect("probe shapes are valid"))
566 .collect();
567
568 assert_eq!(
569 verified_kernels(cands, &probes, run),
570 cands.to_vec(),
571 "the real kernels must all verify on this CPU"
572 );
573
574 let liar = cands[cands.len() - 1];
575 let kept = verified_kernels(cands, &probes, |k, lut, args| {
576 let s = run(k, lut, args);
577 if k == liar {
578 s + 1.0
579 } else {
580 s
581 }
582 });
583 assert!(!kept.contains(&liar), "{liar} lied and was still accepted");
584 assert_eq!(kept.len(), cands.len() - 1, "only the liar should be dropped");
585
586 let none = verified_kernels(cands, &probes, |_, _, _| f32::NAN);
587 assert!(none.is_empty(), "every kernel lied but {none:?} survived");
588 }
589
590 #[test]
593 fn calibration_picks_a_supported_kernel() {
594 let k = calibrated();
595 assert!(
596 cpu_kernels().contains(&k) || matches!(k, Kernel::Scalar(_)),
597 "calibration returned {k}, which is not executable here"
598 );
599 assert_ne!(
600 k,
601 Kernel::Scalar(ScalarReason::SelfCheckFailed),
602 "a kernel this CPU claims to support disagreed with the scalar reference"
603 );
604 }
605
606 #[test]
611 fn every_supported_kernel_matches_scalar_bitwise() {
612 let kernels = kernels_under_test();
613 for seed in [0x452821E638D01377u64, 0x13198A2E03707344, 0xBE5466CF34E90C6C] {
617 check_shapes(&kernels, Rng(seed));
618 }
619 eprintln!("kernels exercised: {kernels:?}");
620 }
621
622 fn check_shapes(kernels: &[Kernel], mut rng: Rng) {
623 for &nq in &[1usize, 3, 7, 8, 9, 16, 17, 32] {
624 for &nbits in &[1usize, 2, 4] {
625 for &dim in &[8usize, 16, 40, 48, 96, 128, 200, 256] {
626 let p = ColbertPacking::new(nbits).unwrap();
627 let n = 1usize << nbits;
628 let mut w: Vec<f32> = (0..n).map(|_| rng.f32(-0.4, 0.4)).collect();
629 w.sort_by(|a, b| a.total_cmp(b));
630 let lut = Lut::new(&p, &w).unwrap();
631 let query: Vec<f32> = (0..nq * dim).map(|_| rng.f32(-1.0, 1.0)).collect();
632 let q = PreparedQuery::new(&lut, &query, nq, dim).unwrap();
633 let ntok = 11;
634 let pdim = dim / lut.keys_per_byte();
635 let row_stride = pdim + 3;
636 let packed: Vec<u8> = (0..ntok * row_stride).map(|_| (rng.next() >> 56) as u8).collect();
637 let ncent = 5;
638 let codes: Vec<u32> = (0..ntok).map(|_| (rng.next() % ncent as u64) as u32).collect();
639 let cdot: Vec<f32> = (0..ncent * nq).map(|_| rng.f32(-1.0, 1.0)).collect();
640 let inv: Vec<f32> = (0..ntok).map(|_| rng.f32(0.5, 1.5)).collect();
641 for with_cdot in [false, true] {
642 for with_inv in [false, true] {
643 let args = Args {
644 query: &q,
645 packed: &packed,
646 row_stride,
647 n_tokens: ntok,
648 codes: if with_cdot {
649 Codes::U32(&codes)
650 } else {
651 Codes::None
652 },
653 cdot: if with_cdot { &cdot } else { q.zeros() },
654 cdot_stride: if with_cdot { nq } else { 0 },
655 inv_norms: if with_inv { Some(&inv) } else { None },
656 };
657 let want = scalar(&lut, &args);
658 for &k in kernels {
659 let got = run(k, &lut, &args);
660 assert_eq!(
661 got.to_bits(),
662 want.to_bits(),
663 "{k}: nq {nq} nbits {nbits} dim {dim} cdot {with_cdot} inv {with_inv}: {got} vs {want}"
664 );
665 }
666 }
667 }
668 }
669 }
670 }
671 }
672}