Skip to main content

search_bench/
search_bench.rs

1//! Search latency on a saved index, without Python.
2//!
3//! cargo run --release --example search_bench -- <index.rvec> <queries.f32> <truth.u32> [ef,...]
4//!
5//! `queries.f32` holds row-major float32 query vectors; `truth.u32` the ids
6//! of the 10 true nearest neighbors of each query (record ids are their row
7//! numbers as strings, as written by bench/ann_bench.py).
8
9use std::time::Instant;
10
11use recern_vector::{Database, SearchOptions};
12
13fn read<T: Copy>(path: &str, from: fn([u8; 4]) -> T) -> Vec<T> {
14    std::fs::read(path)
15        .expect(path)
16        .chunks_exact(4)
17        .map(|b| from(b.try_into().unwrap()))
18        .collect()
19}
20
21fn main() {
22    let args: Vec<String> = std::env::args().skip(1).collect();
23    let db = Database::open(&args[0]).expect("index");
24    let c = db.collections().next().expect("a collection");
25    let dim = c.config().dim;
26    let queries = read(&args[1], f32::from_le_bytes);
27    let truth = read(&args[2], u32::from_le_bytes);
28    let efs: Vec<usize> = args
29        .get(3)
30        .map_or("10,20,40,80,160", String::as_str)
31        .split(',')
32        .map(|e| e.parse().unwrap())
33        .collect();
34    let n = queries.len() / dim;
35
36    println!(
37        "{:>5}  {:>8}  {:>9}  {:>9}  {:>10}",
38        "ef", "recall", "mean µs", "p50 µs", "distances"
39    );
40    for ef in efs {
41        let options = SearchOptions::default().ef(ef);
42        let (mut times, mut found, mut distances) = (Vec::with_capacity(n), 0, 0);
43        for (q, t) in queries.chunks_exact(dim).zip(truth.chunks_exact(10)) {
44            let start = Instant::now();
45            let report = c.explain(q, 10, &options).unwrap();
46            times.push(start.elapsed().as_secs_f64() * 1e6);
47            distances += report.distance_computations;
48            found += report
49                .hits
50                .iter()
51                .filter(|h| t.contains(&h.id.parse::<u32>().unwrap()))
52                .count();
53        }
54        let mean = times.iter().sum::<f64>() / n as f64;
55        times.sort_by(f64::total_cmp);
56        println!(
57            "{ef:>5}  {:>8.4}  {mean:>9.1}  {:>9.1}  {:>10}",
58            found as f64 / (n * 10) as f64,
59            times[n / 2],
60            distances / n
61        );
62    }
63}