Skip to main content

benchmark/
benchmark.rs

1use git_vdb::{CollectionConfig, Database, IndexConfig, Point, Query, QueryParams};
2use serde::Serialize;
3use std::collections::BTreeSet;
4use std::env;
5use std::process::Command;
6use std::time::Instant;
7
8const SEED: u64 = 0x6769_7476_6462_626d;
9
10#[derive(Serialize)]
11struct Report {
12    seed: u64,
13    points: usize,
14    dimension: usize,
15    queries: usize,
16    clusters: usize,
17    index: IndexConfig,
18    root: String,
19    build_ms: u128,
20    exact_query_ms: u128,
21    approximate_query_ms: u128,
22    recall_at_1: f64,
23    recall_at_5: f64,
24    recall_at_10: f64,
25    median_scored_fraction: f64,
26    loose_objects: usize,
27    loose_kib: usize,
28    revision: String,
29    target: String,
30}
31
32fn main() -> git_vdb::Result<()> {
33    let args: Vec<_> = env::args().collect();
34    let count = argument(&args, 1, 1_000);
35    let dimension = argument(&args, 2, 768);
36    let query_count = argument(&args, 3, 100);
37    let clusters = 32.min(count.max(1));
38    let mut rng = SplitMix64(SEED);
39    let centers: Vec<Vec<f32>> = (0..clusters)
40        .map(|_| (0..dimension).map(|_| rng.signed()).collect())
41        .collect();
42    let mut points = Vec::with_capacity(count);
43    for id in 0..count {
44        let center = &centers[id % clusters];
45        let vector = center
46            .iter()
47            .map(|component| component + rng.signed() * 0.08)
48            .collect();
49        points.push(Point {
50            id: (id as u64).into(),
51            vector,
52            payload: Default::default(),
53        });
54    }
55    let queries: Vec<Vec<f32>> = (0..query_count)
56        .map(|index| {
57            centers[index % clusters]
58                .iter()
59                .map(|component| component + rng.signed() * 0.04)
60                .collect()
61        })
62        .collect();
63
64    let temp = tempfile::TempDir::new().expect("temporary benchmark repository");
65    let db = Database::init_bare(temp.path())?;
66    let config = CollectionConfig {
67        dimension,
68        ..CollectionConfig::default()
69    };
70    let collection = db.create_collection("benchmark", config.clone())?;
71    let started = Instant::now();
72    let root = collection.upsert(points)?.root;
73    let build_ms = started.elapsed().as_millis();
74
75    let exact_started = Instant::now();
76    let mut exact = Vec::new();
77    for vector in &queries {
78        exact.push(collection.query(Query {
79            vector: vector.clone(),
80            limit: 10,
81            params: QueryParams {
82                exact: Some(true),
83                ..QueryParams::default()
84            },
85            ..Query::default()
86        })?);
87    }
88    let exact_query_ms = exact_started.elapsed().as_millis();
89
90    let approximate_started = Instant::now();
91    let mut approximate = Vec::new();
92    for vector in &queries {
93        approximate.push(collection.query(Query {
94            vector: vector.clone(),
95            limit: 10,
96            params: QueryParams {
97                exact: Some(false),
98                ..QueryParams::default()
99            },
100            ..Query::default()
101        })?);
102    }
103    let approximate_query_ms = approximate_started.elapsed().as_millis();
104
105    let recall = |k: usize| -> f64 {
106        exact
107            .iter()
108            .zip(&approximate)
109            .map(|(oracle, result)| {
110                let wanted: BTreeSet<_> = oracle.points.iter().take(k).map(|p| &p.id).collect();
111                result
112                    .points
113                    .iter()
114                    .take(k)
115                    .filter(|p| wanted.contains(&p.id))
116                    .count() as f64
117                    / k as f64
118            })
119            .sum::<f64>()
120            / query_count.max(1) as f64
121    };
122    let mut fractions: Vec<_> = approximate
123        .iter()
124        .map(|result| result.stats.vectors_scored as f64 / count.max(1) as f64)
125        .collect();
126    fractions.sort_by(f64::total_cmp);
127    let median_scored_fraction = fractions.get(fractions.len() / 2).copied().unwrap_or(0.0);
128    let count_objects = Command::new("git")
129        .arg("--git-dir")
130        .arg(temp.path())
131        .args(["count-objects", "-v"])
132        .output()
133        .expect("git count-objects");
134    let count_objects = String::from_utf8(count_objects.stdout).expect("UTF-8 Git output");
135    let metric = |name: &str| {
136        count_objects
137            .lines()
138            .find_map(|line| line.strip_prefix(&format!("{name}: ")))
139            .and_then(|value| value.parse().ok())
140            .unwrap_or(0)
141    };
142    let revision = Command::new("git")
143        .args(["rev-parse", "--short=12", "HEAD"])
144        .output()
145        .ok()
146        .filter(|output| output.status.success())
147        .and_then(|output| String::from_utf8(output.stdout).ok())
148        .map(|value| value.trim().to_owned())
149        .unwrap_or_else(|| "uncommitted".into());
150    println!(
151        "{}",
152        serde_json::to_string_pretty(&Report {
153            seed: SEED,
154            points: count,
155            dimension,
156            queries: query_count,
157            clusters,
158            index: config.index,
159            root: root.0,
160            build_ms,
161            exact_query_ms,
162            approximate_query_ms,
163            recall_at_1: recall(1),
164            recall_at_5: recall(5),
165            recall_at_10: recall(10),
166            median_scored_fraction,
167            loose_objects: metric("count"),
168            loose_kib: metric("size"),
169            revision,
170            target: format!("{}-{}", env::consts::OS, env::consts::ARCH),
171        })?
172    );
173    Ok(())
174}
175
176fn argument(args: &[String], index: usize, default: usize) -> usize {
177    args.get(index)
178        .map(|value| value.parse().expect("benchmark arguments must be integers"))
179        .unwrap_or(default)
180}
181
182struct SplitMix64(u64);
183
184impl SplitMix64 {
185    fn next(&mut self) -> u64 {
186        self.0 = self.0.wrapping_add(0x9e37_79b9_7f4a_7c15);
187        let mut value = self.0;
188        value = (value ^ (value >> 30)).wrapping_mul(0xbf58_476d_1ce4_e5b9);
189        value = (value ^ (value >> 27)).wrapping_mul(0x94d0_49bb_1331_11eb);
190        value ^ (value >> 31)
191    }
192
193    fn signed(&mut self) -> f32 {
194        let unit = (self.next() >> 40) as f32 / (1_u64 << 24) as f32;
195        unit * 2.0 - 1.0
196    }
197}