Skip to main content

maxsim_lut/kernel/
mod.rs

1//! Kernel dispatch and the scalar reference.
2//!
3//! Every SIMD kernel computes the identical integer accumulator per (query
4//! row, doc token) and applies the identical float epilogue expression
5//! `(sqw[q] · acc as f32 + crow[q]) · inv`, so all paths return the same
6//! bits. The integration tests pin this.
7//!
8//! # Choosing among the kernels of one architecture
9//!
10//! Both supported architectures now carry more than one int8 dot
11//! instruction, and which is fastest is a property of the core, not of the
12//! instruction set:
13//!
14//! | | wide instruction | narrow instruction | who wins |
15//! |---|---|---|---|
16//! | aarch64 | `smmla`, 32 MACs | `sdot`, 16 MACs | Neoverse N2 issues both per cycle, so `smmla` doubles throughput; Apple M-series issues `smmla` at half rate and loses |
17//! | x86_64 | `vpdpbusd` zmm, 64 MACs | `vpdpbusd` ymm / the AVX2 triple, 32 MACs | depends on how the part implements 512-bit ops and on how far it downclocks |
18//!
19//! Feature detection alone therefore cannot pick the fastest path: two cores
20//! with identical feature bits disagree. [`select`] instead **measures**
21//! them. On the first dispatch of a process, [`calibrated`] scores a small
22//! synthetic document with each candidate, checks it against the scalar
23//! reference, times the survivors interleaved, and caches the winner.
24//!
25//! This is safe to do at runtime precisely because the kernels are
26//! bit-identical: calibration can change how long a search takes, never what
27//! it returns. It costs a few hundred microseconds, once. Set
28//! `MAXSIM_LUT_NO_CALIBRATE=1` to take the first listed kernel instead, or
29//! `MAXSIM_LUT_KERNEL=<name>` to pin one.
30
31use 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/// Which code path scores a given shape on this CPU. Returned by
52/// [`Lut::kernel`]; the `Display` form is meant for benchmark output.
53#[derive(Debug, Clone, Copy, PartialEq, Eq)]
54pub enum Kernel {
55    /// The scalar reference, with the reason the SIMD paths were not taken.
56    Scalar(ScalarReason),
57    /// aarch64 NEON: `tbl` expansion, `sdot` accumulation.
58    NeonSdot,
59    /// aarch64 NEON + I8MM: `tbl` expansion, `smmla` 2×2 matrix accumulation
60    /// (32 MACs per instruction). Faster than `sdot` on cores that issue
61    /// both at the same rate (Arm Neoverse N2/V1/V2), not on Apple cores.
62    NeonI8mm,
63    /// x86_64 AVX2: `pshufb` expansion, `maddubs`/`madd` accumulation.
64    Avx2,
65    /// x86_64 AVX-VNNI: the AVX2 tile shape with 256-bit `vpdpbusd`, for
66    /// cores that have VNNI without AVX-512 (Alder Lake and later hybrids).
67    Avx2Vnni,
68    /// x86_64 AVX-512 with VNNI: `pshufb` expansion, `vpdpbusd` accumulation.
69    Avx512Vnni,
70}
71
72/// Why the scalar path runs.
73#[derive(Debug, Clone, Copy, PartialEq, Eq)]
74pub enum ScalarReason {
75    /// [`Lut::force_scalar`] or `MAXSIM_LUT_FORCE_SCALAR=1`.
76    Forced,
77    /// `dim % 8 != 0` or `dim > MAX_DIM`.
78    DimNotSimdAligned,
79    /// The packing could not be factored into nibble tables (nbits 8).
80    NoNibbleTables,
81    /// This CPU lacks the required feature (NEON `dotprod`, AVX2), or the
82    /// architecture has no SIMD path at all.
83    CpuUnsupported,
84    /// Every SIMD kernel this CPU claims to support disagreed with the
85    /// scalar reference on the calibration probe, so none was trusted. This
86    /// should be unreachable; it means a kernel is miscompiled or the CPU
87    /// misreports a feature, and it trades speed for a correct answer.
88    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    /// `true` for any SIMD path.
106    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
120/// `MAXSIM_LUT_KERNEL=<name>` pins a SIMD kernel by its `Display` name
121/// (`neon-sdot`, `neon-i8mm`, `avx2`, `avx2-vnni`, `avx512-vnni`), for
122/// benchmarking a path calibration would not choose. It is honoured only
123/// when the CPU supports that kernel and the shape is SIMD eligible;
124/// otherwise dispatch proceeds normally. Read once per process.
125fn 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
137/// Every kernel this CPU can execute, in the order to try when calibration
138/// is switched off (widest instruction first, which is the right guess on
139/// most cores). Empty if the CPU has no SIMD path.
140fn 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
171/// `MAXSIM_LUT_NO_CALIBRATE=1` skips the measurement and takes the first
172/// entry of [`cpu_kernels`]. For reproducing a specific dispatch, and for
173/// hosts that cannot spare the one-off cost at startup.
174fn 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
183/// A synthetic index and query the calibrator can score, owning its buffers.
184struct 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    /// `None` only if the shape is invalid, which the call sites' constants
195    /// rule out; it keeps calibration total without a panic.
196    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        // Deterministic pseudo-random content. The kernels never branch on
203        // values, so any well-mixed data times the same and exercises the
204        // same paths.
205        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
234/// Shapes the self-check scores. The first is the ColBERT working point and
235/// is also the one timed; the rest exist to reach the paths a single shape
236/// would miss: an odd query-row count (masked fold tails), a `dim` whose
237/// packed row ends mid-vector (the expansion tail), a row count just past a
238/// block boundary, and each supported code width.
239const PROBE_SHAPES: [(usize, usize, usize, usize); 4] = [
240    // (nbits, dim, n_query_rows, n_doc_tokens)
241    (4, 128, 32, 64),
242    (4, 40, 9, 5),
243    (2, 96, 17, 7),
244    (1, 256, 1, 3),
245];
246
247/// The fastest *verified* kernel for this CPU, decided once per process.
248///
249/// Each candidate is scored against the scalar reference on every
250/// [`PROBE_SHAPES`] entry and dropped if it disagrees anywhere; the
251/// survivors are then timed on the first shape with their arms interleaved
252/// and reduced by minimum, so a scheduling hiccup has to hit every
253/// repetition of one arm to change the verdict.
254///
255/// The check exists because CI cannot own every microarchitecture this crate
256/// emits code for: a kernel whose instruction no test machine has must prove
257/// itself on the host before dispatch will use it. It is a smoke test rather
258/// than a proof, which is why it spans several shapes instead of one.
259fn 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
271/// Candidates that reproduce the scalar reference bit for bit on every
272/// probe. `score` is the kernel runner, taken as an argument so tests can
273/// substitute one that lies.
274fn 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
290/// Verify the candidates, then time the survivors on the production shape.
291///
292/// `fallback` is returned unmeasured if the probes cannot be built. If they
293/// build and *no* candidate reproduces the reference, the result is
294/// `Scalar(SelfCheckFailed)`: a wrong fast answer is worse than a slow right
295/// one.
296fn 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    // A candidate this CPU claims to support and cannot reproduce is a bug in
310    // this crate or a lying CPU; be loud about it where asserts are on.
311    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
344/// The dispatch decision, in one place, so [`Lut::kernel`] and
345/// [`maxsim`] cannot disagree.
346pub(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    // A pin from the host, then from the environment; both are honoured only
357    // for a kernel this CPU can actually execute.
358    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
366/// Every SIMD kernel this CPU can execute, widest instruction first; empty
367/// on a CPU or architecture with no SIMD path. The scalar reference always
368/// runs everywhere and is not listed.
369///
370/// Pass one to [`crate::Lut::pin_kernel`] to bypass calibration.
371pub fn supported_kernels() -> &'static [Kernel] {
372    cpu_kernels()
373}
374
375/// Run the kernel calibration now and return the kernel it chose.
376///
377/// Calibration otherwise happens inside the first [`crate::Scorer::score`]
378/// of the process, which charges one query a few hundred microseconds it did
379/// not expect. A server can call this at startup instead, so the cost lands
380/// before the first request rather than inside it. Calling it more than
381/// once, or from several threads, is harmless: the decision is made once.
382///
383/// The returned kernel is what dispatch will pick *absent* an override; a
384/// [`crate::Lut`] with [`crate::Lut::force_scalar`] or
385/// [`crate::Lut::pin_kernel`] set, a non-SIMD `dim`, or a table with no
386/// nibble factorisation still routes elsewhere. Ask
387/// [`crate::Lut::kernel`] for the decision that applies to a specific table
388/// and dimension.
389pub fn warm_up() -> Kernel {
390    calibrated()
391}
392
393/// Everything a kernel needs, already validated by [`crate::Scorer`].
394pub(crate) struct Args<'a> {
395    pub query: &'a PreparedQuery,
396    /// `[n_tokens · row_stride]` (at least) packed residual bytes.
397    pub packed: &'a [u8],
398    pub row_stride: usize,
399    pub n_tokens: usize,
400    /// Centroid ids; `Codes::None` iff `cdot_stride == 0`.
401    pub codes: Codes<'a>,
402    /// Centroid-major `[num_centroids · nq]` scores, or `[nq]` zeros with stride 0.
403    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        // Range was validated by the scorer; stride 0 selects the zero row.
419        let off = self.codes.id(t) * self.cdot_stride;
420        debug_assert!(off + self.query.n_tokens() <= self.cdot.len());
421        // SAFETY of the later reads: off + nq <= cdot.len() (validated).
422        unsafe { self.cdot.as_ptr().add(off) }
423    }
424}
425
426// Per-thread kernel scratch (`best`, `accs`), reused across the thousands of
427// per-candidate calls of a search. The kernels size-and-initialise it on
428// entry, so no state leaks between calls.
429#[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
435/// Runtime-dispatched MaxSim. `args` must come from `Scorer::validate`.
436pub(crate) fn maxsim(lut: &Lut, args: &Args<'_>) -> f32 {
437    run(select(lut, args.query.dim()), lut, args)
438}
439
440/// Every kernel this CPU can execute, scalar first. Used by the tests so an
441/// AVX-512 machine still exercises the AVX2 path, and a core with both NEON
442/// instructions exercises the one calibration did not pick; `select` alone
443/// would run only the winner.
444#[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
451/// Run a specific kernel. `kernel` must be executable on this CPU and, for
452/// the SIMD variants, `select` must not have returned a `Scalar` reason
453/// other than `Forced` for this `lut`/`dim` (i.e. dim aligned, tables present).
454pub(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            // SAFETY: `select` checked `dotprod`, dim % 8 == 0, dim <= MAX_DIM,
461            // nibble tables present; the scorer validated every slice length.
462            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            // SAFETY: as above, with `i8mm` checked.
468            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            // SAFETY: as above, with AVX2 checked.
474            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            // SAFETY: as above, with avx2 + avxvnni checked.
480            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            // SAFETY: as above, with avx512f/bw/vnni checked.
486            unsafe { avx512::maxsim(lut, args, best, accs) }
487        }),
488        #[allow(unreachable_patterns)]
489        _ => scalar(lut, args),
490    }
491}
492
493/// Scalar reference. Doc-token-outer: expand each stored token's bytes to
494/// int8 weights once, amortised over all query rows.
495pub(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            // SAFETY: crow points at >= nq readable f32s (validated).
522            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    /// The self-check is the only thing standing between an unproven kernel
553    /// and wrong scores, so prove it actually rejects one. A truthful runner
554    /// keeps every candidate; one that corrupts a single kernel's result
555    /// drops exactly that kernel; one that corrupts all of them leaves
556    /// nothing, which is what makes dispatch fall back to the reference.
557    #[test]
558    fn the_self_check_rejects_a_kernel_that_disagrees() {
559        let cands = cpu_kernels();
560        if cands.is_empty() {
561            return; // no SIMD on this target; nothing to verify
562        }
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    /// Calibration only ever returns something dispatch may legally run: a
591    /// kernel this CPU supports, or the scalar reference if none verified.
592    #[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    /// Every executable kernel agrees bitwise with the scalar reference. This
607    /// is the test that reaches AVX2 on an AVX-512 host and AVX-512 on any
608    /// host that has it; the integration tests only see whatever `select`
609    /// picks.
610    #[test]
611    fn every_supported_kernel_matches_scalar_bitwise() {
612        let kernels = kernels_under_test();
613        // Several independent draws per shape: one seed proves a shape works
614        // for one arrangement of bytes, not that the fold and the tails are
615        // right for the values that land near their boundaries.
616        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}