Skip to main content

real_index/
real_index.rs

1//! Score a real next-plaid index with the crate: parity against float
2//! decompression, ns/token on real residual codes, and how loose the exact
3//! per-token skip bound is on real data.
4//!
5//!   cargo run --release --example real_index -- <index_dir> <queries.npy> [n_queries] [ref_docs]
6//!
7//! `index_dir` is a next-plaid ≥1.7 index (`metadata.json`, `centroids.npy`,
8//! `bucket_weights.npy`, `<i>.codes.npy`, `<i>.residuals.npy`,
9//! `<i>.inv_norms.npy`, `doclens.<i>.json`). `queries.npy` is f32
10//! `[n, n_tokens, dim]`. No dependencies: a 60-line NPY reader below.
11
12// The float reference deliberately spells out the index arithmetic.
13#![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
21/// Minimal NPY (v1/v2, C order) reader: returns (shape, dtype descr, raw bytes).
22fn 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
104/// `"key": <int>` out of a small JSON file, no parser.
105fn 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        // Stage-1 product the host would already have: centroid-major [ncent × nq].
196        let t = Instant::now();
197        let mut cdot = vec![0f32; ncent * nq_rows];
198        for c in 0..ncent {
199            let cv = &centroids[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        // Float reference on the first `ref_docs` docs: decompress token =
221        // centroid + bucket weight per dim, exact f32 MaxSim with inv norms.
222        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 = &centroids[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                    // ColBERT packing: key k of byte i is bits (7 - k*nbits ..), bucket bit-reversed.
231                    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        // Exact skip bound: |residual term| <= sqw[r] * 127 * Σ|q̂| = max|w| * ||q_r||_1
259        // (up to quantisation). Simulate a sequential pass over each doc's tokens and
260        // count tokens where every row could have been skipped.
261        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                // Actual residual term magnitude for the record.
277                let cv = &centroids[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}