1#![allow(clippy::needless_range_loop)]
14
15use std::fs;
16use std::path::Path;
17use std::time::Instant;
18
19use maxsim_lut::{Codes, DocView, Lut, PreparedQuery, Scorer};
20
21fn read_npy(path: &Path) -> (Vec<usize>, String, Vec<u8>) {
23 let bytes = fs::read(path).unwrap_or_else(|e| panic!("{}: {e}", path.display()));
24 assert_eq!(&bytes[..6], b"\x93NUMPY", "{}: not an npy file", path.display());
25 let (hlen, hstart) = if bytes[6] == 1 {
26 (u16::from_le_bytes([bytes[8], bytes[9]]) as usize, 10)
27 } else {
28 (
29 u32::from_le_bytes([bytes[8], bytes[9], bytes[10], bytes[11]]) as usize,
30 12,
31 )
32 };
33 let header = std::str::from_utf8(&bytes[hstart..hstart + hlen]).unwrap();
34 let descr = header
35 .split("'descr':")
36 .nth(1)
37 .unwrap()
38 .split('\'')
39 .nth(1)
40 .unwrap()
41 .to_string();
42 assert!(
43 header.contains("'fortran_order': False"),
44 "{}: Fortran order unsupported",
45 path.display()
46 );
47 let shape_str = header
48 .split("'shape':")
49 .nth(1)
50 .unwrap()
51 .split('(')
52 .nth(1)
53 .unwrap()
54 .split(')')
55 .next()
56 .unwrap();
57 let shape: Vec<usize> = shape_str
58 .split(',')
59 .map(|s| s.trim())
60 .filter(|s| !s.is_empty())
61 .map(|s| s.parse().unwrap())
62 .collect();
63 (shape, descr, bytes[hstart + hlen..].to_vec())
64}
65
66fn npy_f32(path: &Path) -> (Vec<usize>, Vec<f32>) {
67 let (shape, descr, raw) = read_npy(path);
68 assert_eq!(descr, "<f4", "{}", path.display());
69 (
70 shape,
71 raw.as_chunks::<4>()
72 .0
73 .iter()
74 .map(|c| f32::from_le_bytes(*c))
75 .collect(),
76 )
77}
78
79fn npy_u8(path: &Path) -> (Vec<usize>, Vec<u8>) {
80 let (shape, descr, raw) = read_npy(path);
81 assert_eq!(descr, "|u1", "{}", path.display());
82 (shape, raw)
83}
84
85fn npy_codes(path: &Path) -> Vec<u32> {
86 let (_, descr, raw) = read_npy(path);
87 match descr.as_str() {
88 "<i8" => raw
89 .as_chunks::<8>()
90 .0
91 .iter()
92 .map(|c| i64::from_le_bytes(*c) as u32)
93 .collect(),
94 "<u4" | "<i4" => raw
95 .as_chunks::<4>()
96 .0
97 .iter()
98 .map(|c| u32::from_le_bytes(*c))
99 .collect(),
100 d => panic!("{}: unsupported code dtype {d}", path.display()),
101 }
102}
103
104fn json_usize(text: &str, key: &str) -> usize {
106 let pat = format!("\"{key}\":");
107 let rest = &text[text.find(&pat).unwrap_or_else(|| panic!("missing {key}")) + pat.len()..];
108 rest.trim_start()
109 .chars()
110 .take_while(|c| c.is_ascii_digit())
111 .collect::<String>()
112 .parse()
113 .unwrap()
114}
115
116fn main() {
117 let args: Vec<String> = std::env::args().skip(1).collect();
118 if args.len() < 2 {
119 eprintln!("usage: real_index <index_dir> <queries.npy> [n_queries=20] [ref_docs=200]");
120 std::process::exit(2);
121 }
122 let dir = Path::new(&args[0]);
123 let n_queries: usize = args.get(2).map(|s| s.parse().unwrap()).unwrap_or(20);
124 let ref_docs: usize = args.get(3).map(|s| s.parse().unwrap()).unwrap_or(200);
125
126 let meta = fs::read_to_string(dir.join("metadata.json")).unwrap();
127 let nbits = json_usize(&meta, "nbits");
128 let num_chunks = json_usize(&meta, "num_chunks");
129 let dim = json_usize(&meta, "embedding_dim");
130 let (cshape, centroids) = npy_f32(&dir.join("centroids.npy"));
131 let ncent = cshape[0];
132 assert_eq!(cshape[1], dim);
133 let (_, weights) = npy_f32(&dir.join("bucket_weights.npy"));
134 assert_eq!(weights.len(), 1 << nbits);
135
136 let mut codes: Vec<u32> = Vec::new();
137 let mut inv: Vec<f32> = Vec::new();
138 let mut packed: Vec<u8> = Vec::new();
139 let mut doclens: Vec<usize> = Vec::new();
140 let pdim = dim * nbits / 8;
141 for c in 0..num_chunks {
142 codes.extend(npy_codes(&dir.join(format!("{c}.codes.npy"))));
143 inv.extend(npy_f32(&dir.join(format!("{c}.inv_norms.npy"))).1);
144 let (rshape, r) = npy_u8(&dir.join(format!("{c}.residuals.npy")));
145 assert_eq!(rshape[1], pdim, "residual row width");
146 packed.extend(r);
147 let dl = fs::read_to_string(dir.join(format!("doclens.{c}.json"))).unwrap();
148 doclens.extend(
149 dl.trim()
150 .trim_matches(['[', ']'])
151 .split(',')
152 .map(|s| s.trim().parse::<usize>().unwrap()),
153 );
154 }
155 let ntok_total: usize = doclens.iter().sum();
156 assert_eq!(codes.len(), ntok_total);
157 assert_eq!(inv.len(), ntok_total);
158 assert_eq!(packed.len(), ntok_total * pdim);
159 let mut offsets = Vec::with_capacity(doclens.len() + 1);
160 offsets.push(0usize);
161 for &l in &doclens {
162 offsets.push(offsets.last().unwrap() + l);
163 }
164
165 let (qshape, queries) = npy_f32(Path::new(&args[1]));
166 let (nq_rows, qdim) = (qshape[1], qshape[2]);
167 assert_eq!(qdim, dim);
168 let n_queries = n_queries.min(qshape[0]);
169
170 let lut = Lut::colbert(nbits, &weights).unwrap();
171 println!(
172 "index: {} docs, {} tokens, dim {dim}, nbits {nbits}, {ncent} centroids, {} B/token packed\nkernel: {}\nqueries: {n_queries} × {nq_rows} rows",
173 doclens.len(),
174 ntok_total,
175 pdim,
176 lut.kernel(dim)
177 );
178 let max_w = weights.iter().fold(0f32, |m, &w| m.max(w.abs()));
179
180 let mut t_cdot = 0.0f64;
181 let mut t_score = 0.0f64;
182 let mut max_abs_err = 0f32;
183 let mut sum_rel_err = 0f64;
184 let mut n_ref = 0usize;
185 let mut skippable = 0usize;
186 let mut skip_seen = 0usize;
187 let mut bound_sum = 0f64;
188 let mut resid_abs_sum = 0f64;
189 let mut resid_n = 0usize;
190
191 for qi in 0..n_queries {
192 let qf = &queries[qi * nq_rows * dim..(qi + 1) * nq_rows * dim];
193 let q = PreparedQuery::new(&lut, qf, nq_rows, dim).unwrap();
194
195 let t = Instant::now();
197 let mut cdot = vec![0f32; ncent * nq_rows];
198 for c in 0..ncent {
199 let cv = ¢roids[c * dim..(c + 1) * dim];
200 for r in 0..nq_rows {
201 let qr = &qf[r * dim..(r + 1) * dim];
202 cdot[c * nq_rows + r] = cv.iter().zip(qr).map(|(a, b)| a * b).sum();
203 }
204 }
205 t_cdot += t.elapsed().as_secs_f64();
206
207 let scorer = Scorer::new(&lut, &q).with_centroid_term(&cdot, ncent).unwrap();
208 let mut scores = vec![0f32; doclens.len()];
209 let t = Instant::now();
210 for (d, s) in scores.iter_mut().enumerate() {
211 let (a, b) = (offsets[d], offsets[d + 1]);
212 *s = scorer.score(
213 DocView::new(&packed[a * pdim..b * pdim], b - a, pdim)
214 .codes(Codes::U32(&codes[a..b]))
215 .inv_norms(&inv[a..b]),
216 );
217 }
218 t_score += t.elapsed().as_secs_f64();
219
220 let mut tok = vec![0f32; dim];
223 for d in 0..ref_docs.min(doclens.len()) {
224 let (a, b) = (offsets[d], offsets[d + 1]);
225 let mut best = vec![f32::NEG_INFINITY; nq_rows];
226 for t in a..b {
227 let cv = ¢roids[codes[t] as usize * dim..(codes[t] as usize + 1) * dim];
228 let row = &packed[t * pdim..(t + 1) * pdim];
229 for dd in 0..dim {
230 let kpb = 8 / nbits;
232 let (i, k) = (dd / kpb, dd % kpb);
233 let shift = 8 - nbits * (k + 1);
234 let seg = (row[i] >> shift) as usize & ((1 << nbits) - 1);
235 let mut bucket = 0usize;
236 for bit in 0..nbits {
237 if seg & (1 << bit) != 0 {
238 bucket |= 1 << (nbits - 1 - bit);
239 }
240 }
241 tok[dd] = cv[dd] + weights[bucket];
242 }
243 for r in 0..nq_rows {
244 let qr = &qf[r * dim..(r + 1) * dim];
245 let s: f32 = tok.iter().zip(qr).map(|(x, y)| x * y).sum::<f32>() * inv[t];
246 if s > best[r] {
247 best[r] = s;
248 }
249 }
250 }
251 let reference: f32 = best.iter().sum();
252 let err = (scores[d] - reference).abs();
253 max_abs_err = max_abs_err.max(err);
254 sum_rel_err += (err / reference.abs().max(1e-6)) as f64;
255 n_ref += 1;
256 }
257
258 let l1: Vec<f32> = (0..nq_rows)
262 .map(|r| qf[r * dim..(r + 1) * dim].iter().map(|x| x.abs()).sum())
263 .collect();
264 let bounds: Vec<f32> = l1.iter().map(|&l| l * max_w).collect();
265 bound_sum += bounds.iter().map(|&b| b as f64).sum::<f64>() / nq_rows as f64;
266 for d in 0..ref_docs.min(doclens.len()) {
267 let (a, b) = (offsets[d], offsets[d + 1]);
268 let mut best = vec![f32::NEG_INFINITY; nq_rows];
269 for t in a..b {
270 let crow = &cdot[codes[t] as usize * nq_rows..(codes[t] as usize + 1) * nq_rows];
271 let can_skip = (0..nq_rows).all(|r| (crow[r] + bounds[r]) * inv[t] <= best[r]);
272 skip_seen += 1;
273 if can_skip {
274 skippable += 1;
275 }
276 let cv = ¢roids[codes[t] as usize * dim..(codes[t] as usize + 1) * dim];
278 let row = &packed[t * pdim..(t + 1) * pdim];
279 let kpb = 8 / nbits;
280 for r in 0..nq_rows {
281 let qr = &qf[r * dim..(r + 1) * dim];
282 let mut resid = 0f32;
283 for dd in 0..dim {
284 let (i, k) = (dd / kpb, dd % kpb);
285 let shift = 8 - nbits * (k + 1);
286 let seg = (row[i] >> shift) as usize & ((1 << nbits) - 1);
287 let mut bucket = 0usize;
288 for bit in 0..nbits {
289 if seg & (1 << bit) != 0 {
290 bucket |= 1 << (nbits - 1 - bit);
291 }
292 }
293 resid += qr[dd] * weights[bucket];
294 }
295 resid_abs_sum += resid.abs() as f64;
296 resid_n += 1;
297 let s =
298 (crow[r] + (cv.iter().zip(qr).map(|(x, y)| x * y).sum::<f32>() - crow[r]) + resid)
299 * inv[t];
300 if s > best[r] {
301 best[r] = s;
302 }
303 }
304 }
305 }
306 }
307
308 let tokens_scored = ntok_total as f64 * n_queries as f64;
309 println!(
310 "\nexhaustive stage-2 over all docs: {:.2} ms/query, {:.2} ns/token (kernel only; stage-1 cdot GEMM excluded: {:.1} ms/query naive)",
311 t_score / n_queries as f64 * 1e3,
312 t_score / tokens_scored * 1e9,
313 t_cdot / n_queries as f64 * 1e3
314 );
315 println!(
316 "parity vs float decompression on {n_ref} (query, doc) pairs: max |Δscore| {max_abs_err:.4}, mean rel {:.2e}",
317 sum_rel_err / n_ref.max(1) as f64
318 );
319 println!(
320 "exact skip bound: mean bound {:.3} vs mean |residual term| {:.4}; tokens skippable {}/{} ({:.2}%)",
321 bound_sum / n_queries as f64,
322 resid_abs_sum / resid_n.max(1) as f64,
323 skippable,
324 skip_seen,
325 100.0 * skippable as f64 / skip_seen.max(1) as f64
326 );
327}