#![cfg(test)]
use crate::dist::sqdist;
use crate::partition::{Partitions, Tuning, Vectors};
use crate::rabitq::Bits;
use std::time::Instant;
use yo_common::Rng;
fn from_env(name: &str, default: usize) -> usize {
std::env::var(name)
.ok()
.and_then(|v| v.trim().parse().ok())
.unwrap_or(default)
}
fn probe() -> usize {
from_env("YO_PROBE", 32)
}
fn posting() -> usize {
from_env("YO_POSTING", Tuning::default().posting)
}
const QUERIES: usize = 200;
struct Base {
dim: usize,
data: Vec<f32>,
}
impl Vectors for Base {
fn get(&self, id: u64, into: &mut [f32]) -> bool {
let at = id as usize * self.dim;
let Some(v) = self.data.get(at..at + self.dim) else {
return false;
};
into.copy_from_slice(v);
true
}
}
fn generated(dim: usize, n: usize, clusters: usize, seed: u64) -> Base {
let mut rng = Rng::new(seed);
let draw = |rng: &mut Rng| -> Vec<f32> {
(0..dim)
.map(|i| {
let u = (rng.next_u64() >> 40) as f32 / (1u32 << 24) as f32;
let heavy = if i < dim / 16 { 6.0 } else { 1.0 };
(u * 2.0 - 1.0) * heavy
})
.collect()
};
let centres: Vec<Vec<f32>> = (0..clusters).map(|_| draw(&mut rng)).collect();
let mut data = Vec::with_capacity(n * dim);
for i in 0..n {
let off = draw(&mut rng);
for (c, o) in centres[i % clusters].iter().zip(&off) {
data.push(c + o * 0.7);
}
}
Base { dim, data }
}
fn build(base: &Base) -> Partitions {
let n = base.data.len() / base.dim;
let tuning = Tuning {
probe: probe(),
posting: posting(),
..Tuning::default()
};
let mut ix = Partitions::new(base.dim, Bits::One, 0x51f7, tuning);
let mut buf = vec![0f32; base.dim];
for id in 0..n as u64 {
base.get(id, &mut buf);
ix.insert(id, &buf);
if ix.needs_maintenance() {
ix.maintain(base, 4);
}
}
ix
}
#[test]
#[ignore = "prints a table rather than asserting anything"]
fn where_a_query_spends_its_time() {
let bases: Vec<Base> = match std::env::var("YO_DATASET") {
Ok(dir) => vec![from_dataset(dir.trim())],
Err(_) => vec![
generated(768, 60_000, 24, 5),
generated(1024, 200_000, 24, 5),
],
};
for base in &bases {
let dim = base.dim;
let n = base.data.len() / dim;
let at = Instant::now();
let ix = build(base);
let parts = ix.partitions();
let held: usize = (0..parts).map(|p| ix.posting_parts(p).0.len()).sum();
println!(
"{dim} dims, {n} vectors, {parts} partitions, {held} entries, built in {:?}",
at.elapsed()
);
let queries: Vec<&[f32]> = (0..QUERIES)
.map(|i| {
let at = (i * 7919 % n) * dim;
&base.data[at..at + dim]
})
.collect();
let each = |d: std::time::Duration| d.as_nanos() as f64 / QUERIES as f64 / 1000.0;
let mut sink = 0f32;
let mut order = Vec::new();
let at = Instant::now();
for q in &queries {
ix.probe_order(q, &mut order);
sink += order[0] as f32;
}
let rank = each(at.elapsed());
let centroids = ix.all_centroids();
let at = Instant::now();
for q in &queries {
ix.probe_order(q, &mut order);
let u = ix.quantizer().rotate(q);
for &p in order.iter().take(probe()) {
let prepared = ix
.quantizer()
.query_rotated(&u, ¢roids[p * dim..(p + 1) * dim]);
let (_, _, codes, meta, _) = ix.posting_parts(p);
sink += prepared.distance(&codes[..ix.quantizer().code_bytes()], &meta[0]);
}
}
let prep = each(at.elapsed());
let mut scores: Vec<f32> = Vec::new();
let at = Instant::now();
for q in &queries {
ix.probe_order(q, &mut order);
let u = ix.quantizer().rotate(q);
for &p in order.iter().take(probe()) {
let prepared = ix
.quantizer()
.query_rotated(&u, ¢roids[p * dim..(p + 1) * dim]);
let (_, _, codes, meta, _) = ix.posting_parts(p);
if scores.len() < meta.len() {
scores.resize(meta.len(), 0.0);
}
prepared.scan(codes, meta, &mut scores[..meta.len()]);
sink += scores[0];
}
}
let scan = each(at.elapsed());
let members = probe() * held / parts;
println!(" ranking centroids {rank:8.1} us");
println!(
" preparing the query {:8.1} us ({:.1} us a partition)",
prep - rank,
(prep - rank) / probe() as f64
);
println!(
" scanning the postings {:6.1} us ({:.1} ns a member over {members})",
scan - prep,
(scan - prep) * 1000.0 / members as f64
);
println!(" and the whole search {:6.1} us ({sink})", {
let at = Instant::now();
for q in &queries {
sink += ix.candidates(q, 10)[0].1;
}
each(at.elapsed())
});
}
}
fn from_dataset(dir: &str) -> Base {
let set = dir
.trim_end_matches(['/', '\\'])
.rsplit(['/', '\\'])
.next()
.unwrap_or(dir);
let path = format!("{dir}/{set}_base.fvecs");
let bytes = std::fs::read(&path).unwrap_or_else(|e| panic!("reading {path}: {e}"));
assert!(bytes.len() >= 4, "{path} is too short to hold a vector");
let dim = i32::from_le_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]) as usize;
let stride = 4 + dim * 4;
assert_eq!(bytes.len() % stride, 0, "{path} is not whole vectors");
let mut data = Vec::with_capacity(bytes.len() / stride * dim);
for v in 0..bytes.len() / stride {
for d in 0..dim {
let at = v * stride + 4 + d * 4;
data.push(f32::from_le_bytes([
bytes[at],
bytes[at + 1],
bytes[at + 2],
bytes[at + 3],
]));
}
}
Base { dim, data }
}
#[test]
#[ignore = "prints a number rather than asserting anything"]
fn what_one_centroid_costs() {
for &dim in &[128usize, 768, 1024] {
let base = generated(dim, 4000, 24, 9);
let n = base.data.len() / dim;
let q = &base.data[..dim];
let at = Instant::now();
let mut sink = 0f32;
for _ in 0..100 {
for p in 0..n {
sink += sqdist(q, &base.data[p * dim..(p + 1) * dim]);
}
}
let each = at.elapsed().as_nanos() as f64 / 100.0 / n as f64;
println!("{dim:5} dimensions {each:6.1} ns a centroid ({sink})");
}
}