use std::hint::black_box;
use std::sync::OnceLock;
use std::time::Instant;
use crate::lut::Lut;
use crate::query::PreparedQuery;
use crate::scorer::Codes;
use crate::MAX_DIM;
#[cfg(target_arch = "x86_64")]
mod avx2;
#[cfg(target_arch = "x86_64")]
mod avx2_vnni;
#[cfg(target_arch = "x86_64")]
mod avx512;
#[cfg(target_arch = "aarch64")]
mod neon;
#[cfg(target_arch = "aarch64")]
mod neon_i8mm;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Kernel {
Scalar(ScalarReason),
NeonSdot,
NeonI8mm,
Avx2,
Avx2Vnni,
Avx512Vnni,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ScalarReason {
Forced,
DimNotSimdAligned,
NoNibbleTables,
CpuUnsupported,
SelfCheckFailed,
}
impl std::fmt::Display for Kernel {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Kernel::Scalar(r) => write!(f, "scalar ({r:?})"),
Kernel::NeonSdot => write!(f, "neon-sdot"),
Kernel::NeonI8mm => write!(f, "neon-i8mm"),
Kernel::Avx2 => write!(f, "avx2"),
Kernel::Avx2Vnni => write!(f, "avx2-vnni"),
Kernel::Avx512Vnni => write!(f, "avx512-vnni"),
}
}
}
impl Kernel {
pub fn is_simd(&self) -> bool {
!matches!(self, Kernel::Scalar(_))
}
}
fn env_force_scalar() -> bool {
static FORCE: OnceLock<bool> = OnceLock::new();
*FORCE.get_or_init(|| {
std::env::var("MAXSIM_LUT_FORCE_SCALAR")
.map(|v| v == "1" || v.eq_ignore_ascii_case("true"))
.unwrap_or(false)
})
}
fn env_pin() -> Option<Kernel> {
static PIN: OnceLock<Option<Kernel>> = OnceLock::new();
*PIN.get_or_init(|| match std::env::var("MAXSIM_LUT_KERNEL").ok()?.as_str() {
"neon-sdot" => Some(Kernel::NeonSdot),
"neon-i8mm" => Some(Kernel::NeonI8mm),
"avx2" => Some(Kernel::Avx2),
"avx2-vnni" => Some(Kernel::Avx2Vnni),
"avx512-vnni" => Some(Kernel::Avx512Vnni),
_ => None,
})
}
fn cpu_kernels() -> &'static [Kernel] {
#[cfg(target_arch = "aarch64")]
{
let dotprod = std::arch::is_aarch64_feature_detected!("dotprod");
let i8mm = std::arch::is_aarch64_feature_detected!("i8mm");
match (dotprod, i8mm) {
(true, true) => &[Kernel::NeonI8mm, Kernel::NeonSdot],
(true, false) => &[Kernel::NeonSdot],
(false, true) => &[Kernel::NeonI8mm],
(false, false) => &[],
}
}
#[cfg(target_arch = "x86_64")]
{
let avx512 = is_x86_feature_detected!("avx512f")
&& is_x86_feature_detected!("avx512bw")
&& is_x86_feature_detected!("avx512vnni");
let vnni256 = is_x86_feature_detected!("avxvnni");
let avx2 = is_x86_feature_detected!("avx2");
match (avx512, vnni256, avx2) {
(true, true, _) => &[Kernel::Avx512Vnni, Kernel::Avx2Vnni, Kernel::Avx2],
(true, false, _) => &[Kernel::Avx512Vnni, Kernel::Avx2],
(false, true, _) => &[Kernel::Avx2Vnni, Kernel::Avx2],
(false, false, true) => &[Kernel::Avx2],
(false, false, false) => &[],
}
}
#[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
&[]
}
fn env_no_calibrate() -> bool {
static SKIP: OnceLock<bool> = OnceLock::new();
*SKIP.get_or_init(|| {
std::env::var("MAXSIM_LUT_NO_CALIBRATE")
.map(|v| v == "1" || v.eq_ignore_ascii_case("true"))
.unwrap_or(false)
})
}
struct Probe {
lut: Lut,
query: PreparedQuery,
packed: Vec<u8>,
inv: Vec<f32>,
row_stride: usize,
n_tokens: usize,
}
impl Probe {
fn new(nbits: usize, dim: usize, nq: usize, n_tokens: usize) -> Option<Self> {
let n = 1usize << nbits;
let weights: Vec<f32> = (0..n)
.map(|i| -0.35 + 0.7 * (i as f32 + 0.5) / n as f32)
.collect();
let lut = Lut::colbert(nbits, &weights).ok()?;
let q: Vec<f32> = (0..nq * dim)
.map(|i| ((i * 37 % 251) as f32 / 251.0) - 0.5)
.collect();
let query = PreparedQuery::new(&lut, &q, nq, dim).ok()?;
let row_stride = dim / lut.keys_per_byte();
Some(Self {
packed: (0..n_tokens * row_stride).map(|i| (i * 97 % 256) as u8).collect(),
inv: (0..n_tokens).map(|i| 0.8 + (i % 7) as f32 * 0.05).collect(),
lut,
query,
row_stride,
n_tokens,
})
}
fn args(&self) -> Args<'_> {
Args {
query: &self.query,
packed: &self.packed,
row_stride: self.row_stride,
n_tokens: self.n_tokens,
codes: Codes::None,
cdot: self.query.zeros(),
cdot_stride: 0,
inv_norms: Some(&self.inv),
}
}
}
const PROBE_SHAPES: [(usize, usize, usize, usize); 4] = [
(4, 128, 32, 64),
(4, 40, 9, 5),
(2, 96, 17, 7),
(1, 256, 1, 3),
];
fn calibrated() -> Kernel {
static CHOICE: OnceLock<Kernel> = OnceLock::new();
*CHOICE.get_or_init(|| {
let cands = cpu_kernels();
match cands.first() {
None => Kernel::Scalar(ScalarReason::CpuUnsupported),
Some(&first) if env_no_calibrate() => first,
Some(&first) => measure_fastest(cands, first),
}
})
}
fn verified_kernels<F>(cands: &[Kernel], probes: &[Probe], mut score: F) -> Vec<Kernel>
where
F: FnMut(Kernel, &Lut, &Args<'_>) -> f32,
{
cands
.iter()
.copied()
.filter(|&k| {
probes.iter().all(|p| {
let args = p.args();
score(k, &p.lut, &args).to_bits() == scalar(&p.lut, &args).to_bits()
})
})
.collect()
}
fn measure_fastest(cands: &[Kernel], fallback: Kernel) -> Kernel {
const REPS: usize = 5;
const ITERS: usize = 8;
let mut probes = Vec::with_capacity(PROBE_SHAPES.len());
for (nbits, dim, nq, ntok) in PROBE_SHAPES {
match Probe::new(nbits, dim, nq, ntok) {
Some(p) => probes.push(p),
None => return fallback,
}
}
let verified = verified_kernels(cands, &probes, run);
debug_assert_eq!(
verified.len(),
cands.len(),
"a supported kernel disagreed with the scalar reference: kept {verified:?} of {cands:?}"
);
let Some((&first, rest)) = verified.split_first() else {
return Kernel::Scalar(ScalarReason::SelfCheckFailed);
};
if rest.is_empty() {
return first;
}
let timed = &probes[0];
let args = timed.args();
let mut best = vec![f64::INFINITY; verified.len()];
for _ in 0..REPS {
for (slot, &k) in best.iter_mut().zip(&verified) {
let t = Instant::now();
for _ in 0..ITERS {
black_box(run(k, &timed.lut, &args));
}
*slot = slot.min(t.elapsed().as_secs_f64());
}
}
let mut winner = 0usize;
for i in 1..verified.len() {
if best[i] < best[winner] {
winner = i;
}
}
verified[winner]
}
pub(crate) fn select(lut: &Lut, dim: usize) -> Kernel {
if lut.force_scalar_set() || env_force_scalar() {
return Kernel::Scalar(ScalarReason::Forced);
}
if !dim.is_multiple_of(8) || dim > MAX_DIM {
return Kernel::Scalar(ScalarReason::DimNotSimdAligned);
}
if lut.nibble_tables().is_none() {
return Kernel::Scalar(ScalarReason::NoNibbleTables);
}
for pin in [lut.pinned_kernel(), env_pin()].into_iter().flatten() {
if cpu_kernels().contains(&pin) {
return pin;
}
}
calibrated()
}
pub fn supported_kernels() -> &'static [Kernel] {
cpu_kernels()
}
pub fn warm_up() -> Kernel {
calibrated()
}
pub(crate) struct Args<'a> {
pub query: &'a PreparedQuery,
pub packed: &'a [u8],
pub row_stride: usize,
pub n_tokens: usize,
pub codes: Codes<'a>,
pub cdot: &'a [f32],
pub cdot_stride: usize,
pub inv_norms: Option<&'a [f32]>,
}
impl Args<'_> {
#[inline(always)]
pub(crate) fn inv(&self, t: usize) -> f32 {
match self.inv_norms {
Some(v) => v[t],
None => 1.0,
}
}
#[inline(always)]
pub(crate) fn crow(&self, t: usize) -> *const f32 {
let off = self.codes.id(t) * self.cdot_stride;
debug_assert!(off + self.query.n_tokens() <= self.cdot.len());
unsafe { self.cdot.as_ptr().add(off) }
}
}
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
thread_local! {
static SCRATCH: std::cell::RefCell<(Vec<f32>, Vec<i32>)> =
const { std::cell::RefCell::new((Vec::new(), Vec::new())) };
}
pub(crate) fn maxsim(lut: &Lut, args: &Args<'_>) -> f32 {
run(select(lut, args.query.dim()), lut, args)
}
#[cfg(test)]
fn kernels_under_test() -> Vec<Kernel> {
let mut v = vec![Kernel::Scalar(ScalarReason::Forced)];
v.extend_from_slice(cpu_kernels());
v
}
pub(crate) fn run(kernel: Kernel, lut: &Lut, args: &Args<'_>) -> f32 {
match kernel {
Kernel::Scalar(_) => scalar(lut, args),
#[cfg(target_arch = "aarch64")]
Kernel::NeonSdot => SCRATCH.with(|s| {
let (best, accs) = &mut *s.borrow_mut();
unsafe { neon::maxsim(lut, args, best, accs) }
}),
#[cfg(target_arch = "aarch64")]
Kernel::NeonI8mm => SCRATCH.with(|s| {
let (best, accs) = &mut *s.borrow_mut();
unsafe { neon_i8mm::maxsim(lut, args, best, accs) }
}),
#[cfg(target_arch = "x86_64")]
Kernel::Avx2 => SCRATCH.with(|s| {
let (best, accs) = &mut *s.borrow_mut();
unsafe { avx2::maxsim(lut, args, best, accs) }
}),
#[cfg(target_arch = "x86_64")]
Kernel::Avx2Vnni => SCRATCH.with(|s| {
let (best, accs) = &mut *s.borrow_mut();
unsafe { avx2_vnni::maxsim(lut, args, best, accs) }
}),
#[cfg(target_arch = "x86_64")]
Kernel::Avx512Vnni => SCRATCH.with(|s| {
let (best, accs) = &mut *s.borrow_mut();
unsafe { avx512::maxsim(lut, args, best, accs) }
}),
#[allow(unreachable_patterns)]
_ => scalar(lut, args),
}
}
pub(crate) fn scalar(lut: &Lut, a: &Args<'_>) -> f32 {
let q = a.query;
let nq = q.n_tokens();
let dim = q.dim();
if nq == 0 || a.n_tokens == 0 {
return 0.0;
}
let kpb = lut.keys_per_byte();
let pdim = dim / kpb;
let qv = q.codes();
let sqw = q.sqw();
let mut best = vec![f32::NEG_INFINITY; nq];
let mut w = [0i8; MAX_DIM];
for t in 0..a.n_tokens {
let row = &a.packed[t * a.row_stride..t * a.row_stride + pdim];
for (i, &byte) in row.iter().enumerate() {
w[i * kpb..(i + 1) * kpb].copy_from_slice(lut.expand(byte));
}
let inv = a.inv(t);
let crow = a.crow(t);
for (qi, best_q) in best.iter_mut().enumerate() {
let qrow = &qv[qi * dim..(qi + 1) * dim];
let mut acc = 0i32;
for (qd, wd) in qrow.iter().zip(&w[..dim]) {
acc += *qd as i32 * *wd as i32;
}
let c = unsafe { *crow.add(qi) };
let score = (sqw[qi] * acc as f32 + c) * inv;
if score > *best_q {
*best_q = score;
}
}
}
best.iter().sum()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ColbertPacking;
struct Rng(u64);
impl Rng {
fn next(&mut self) -> u64 {
let mut x = self.0;
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
self.0 = x;
x
}
fn f32(&mut self, lo: f32, hi: f32) -> f32 {
lo + (hi - lo) * ((self.next() >> 40) as f32 / (1u64 << 24) as f32)
}
}
#[test]
fn the_self_check_rejects_a_kernel_that_disagrees() {
let cands = cpu_kernels();
if cands.is_empty() {
return; }
let probes: Vec<Probe> = PROBE_SHAPES
.iter()
.map(|&(nbits, dim, nq, ntok)| Probe::new(nbits, dim, nq, ntok).expect("probe shapes are valid"))
.collect();
assert_eq!(
verified_kernels(cands, &probes, run),
cands.to_vec(),
"the real kernels must all verify on this CPU"
);
let liar = cands[cands.len() - 1];
let kept = verified_kernels(cands, &probes, |k, lut, args| {
let s = run(k, lut, args);
if k == liar {
s + 1.0
} else {
s
}
});
assert!(!kept.contains(&liar), "{liar} lied and was still accepted");
assert_eq!(kept.len(), cands.len() - 1, "only the liar should be dropped");
let none = verified_kernels(cands, &probes, |_, _, _| f32::NAN);
assert!(none.is_empty(), "every kernel lied but {none:?} survived");
}
#[test]
fn calibration_picks_a_supported_kernel() {
let k = calibrated();
assert!(
cpu_kernels().contains(&k) || matches!(k, Kernel::Scalar(_)),
"calibration returned {k}, which is not executable here"
);
assert_ne!(
k,
Kernel::Scalar(ScalarReason::SelfCheckFailed),
"a kernel this CPU claims to support disagreed with the scalar reference"
);
}
#[test]
fn every_supported_kernel_matches_scalar_bitwise() {
let kernels = kernels_under_test();
for seed in [0x452821E638D01377u64, 0x13198A2E03707344, 0xBE5466CF34E90C6C] {
check_shapes(&kernels, Rng(seed));
}
eprintln!("kernels exercised: {kernels:?}");
}
fn check_shapes(kernels: &[Kernel], mut rng: Rng) {
for &nq in &[1usize, 3, 7, 8, 9, 16, 17, 32] {
for &nbits in &[1usize, 2, 4] {
for &dim in &[8usize, 16, 40, 48, 96, 128, 200, 256] {
let p = ColbertPacking::new(nbits).unwrap();
let n = 1usize << nbits;
let mut w: Vec<f32> = (0..n).map(|_| rng.f32(-0.4, 0.4)).collect();
w.sort_by(|a, b| a.total_cmp(b));
let lut = Lut::new(&p, &w).unwrap();
let query: Vec<f32> = (0..nq * dim).map(|_| rng.f32(-1.0, 1.0)).collect();
let q = PreparedQuery::new(&lut, &query, nq, dim).unwrap();
let ntok = 11;
let pdim = dim / lut.keys_per_byte();
let row_stride = pdim + 3;
let packed: Vec<u8> = (0..ntok * row_stride).map(|_| (rng.next() >> 56) as u8).collect();
let ncent = 5;
let codes: Vec<u32> = (0..ntok).map(|_| (rng.next() % ncent as u64) as u32).collect();
let cdot: Vec<f32> = (0..ncent * nq).map(|_| rng.f32(-1.0, 1.0)).collect();
let inv: Vec<f32> = (0..ntok).map(|_| rng.f32(0.5, 1.5)).collect();
for with_cdot in [false, true] {
for with_inv in [false, true] {
let args = Args {
query: &q,
packed: &packed,
row_stride,
n_tokens: ntok,
codes: if with_cdot {
Codes::U32(&codes)
} else {
Codes::None
},
cdot: if with_cdot { &cdot } else { q.zeros() },
cdot_stride: if with_cdot { nq } else { 0 },
inv_norms: if with_inv { Some(&inv) } else { None },
};
let want = scalar(&lut, &args);
for &k in kernels {
let got = run(k, &lut, &args);
assert_eq!(
got.to_bits(),
want.to_bits(),
"{k}: nq {nq} nbits {nbits} dim {dim} cdot {with_cdot} inv {with_inv}: {got} vs {want}"
);
}
}
}
}
}
}
}
}