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 = ¢ers[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}