Skip to main content

Lut

Struct Lut 

Source
pub struct Lut { /* private fields */ }
Expand description

The document-side lookup state for one residual codec: a table turning each packed residual byte directly into its 8/nbits int8 bucket weights, plus the dequantisation scale.

Build once per index (it depends only on the bucket weights and the packing), share across threads.

Implementations§

Source§

impl Lut

Source

pub fn new<P: Packing>( packing: &P, bucket_weights: &[f32], ) -> Result<Self, Error>

Build the table from a packing layout and the codec’s 2^nbits bucket weights (f32, in bucket-index order).

Weights are quantised symmetrically to int8 with scale = max|w| / 127.

Examples found in repository?
examples/bench.rs (line 73)
55fn main() {
56    let args: Vec<usize> = std::env::args()
57        .skip(1)
58        .map(|a| a.parse().expect("integer arg"))
59        .collect();
60    let dim = args.first().copied().unwrap_or(128);
61    let nbits = args.get(1).copied().unwrap_or(4);
62    let nq = args.get(2).copied().unwrap_or(32);
63    let ntok = args.get(3).copied().unwrap_or(240);
64    let ndocs = args.get(4).copied().unwrap_or(1024);
65    let ncent = 16_384;
66    let reps = 9;
67    let mut rng = Rng(0x9E3779B97F4A7C15);
68
69    let p = ColbertPacking::new(nbits).unwrap();
70    let nb = 1usize << nbits;
71    let mut w: Vec<f32> = (0..nb).map(|_| rng.f32(-0.4, 0.4)).collect();
72    w.sort_by(|a, b| a.total_cmp(b));
73    let lut = Lut::new(&p, &w).unwrap();
74
75    let query: Vec<f32> = (0..nq * dim).map(|_| rng.f32(-1.0, 1.0)).collect();
76    let q = PreparedQuery::new(&lut, &query, nq, dim).unwrap();
77    let cdot: Vec<f32> = (0..ncent * nq).map(|_| rng.f32(-1.0, 1.0)).collect();
78
79    let pdim = dim / p.keys_per_byte();
80    let mut packed = vec![0u8; ndocs * ntok * pdim];
81    for b in packed.iter_mut() {
82        *b = (rng.next() >> 56) as u8;
83    }
84    let codes: Vec<u32> = (0..ndocs * ntok)
85        .map(|_| (rng.next() % ncent as u64) as u32)
86        .collect();
87    let inv: Vec<f32> = (0..ndocs * ntok).map(|_| rng.f32(0.8, 1.2)).collect();
88    let docs: Vec<DocView> = (0..ndocs)
89        .map(|d| {
90            DocView::new(&packed[d * ntok * pdim..(d + 1) * ntok * pdim], ntok, pdim)
91                .codes(Codes::U32(&codes[d * ntok..(d + 1) * ntok]))
92                .inv_norms(&inv[d * ntok..(d + 1) * ntok])
93        })
94        .collect();
95
96    // One arm per executable kernel, plus the scalar reference. `dispatch`
97    // is what an unpinned host would get, named so the calibrated choice is
98    // visible next to the kernels it chose between.
99    let dispatched = lut.kernel(dim);
100    let mut arms: Vec<Arm> = Vec::new();
101    for &k in supported_kernels() {
102        arms.push(Arm {
103            label: if k == dispatched {
104                format!("{k} (dispatched)")
105            } else {
106                format!("{k}")
107            },
108            lut: lut.clone().pin_kernel(Some(k)),
109            ns: Vec::new(),
110            checksum: 0.0,
111        });
112    }
113    if !dispatched.is_simd() {
114        arms.push(Arm {
115            label: format!("{dispatched} (dispatched)"),
116            lut: lut.clone(),
117            ns: Vec::new(),
118            checksum: 0.0,
119        });
120    }
121    arms.push(Arm {
122        label: "scalar reference".to_string(),
123        lut: lut.clone().force_scalar(true),
124        ns: Vec::new(),
125        checksum: 0.0,
126    });
127
128    println!(
129        "dim {dim}, nbits {nbits}, {nq} query tokens, {ndocs} docs × {ntok} tokens, {ncent} centroids\narch {}, dispatched kernel: {dispatched}, {reps} interleaved rounds",
130        std::env::consts::ARCH,
131    );
132
133    let mut out = vec![0.0f32; ndocs];
134    for arm in arms.iter_mut() {
135        let s = Scorer::new(&arm.lut, &q)
136            .with_centroid_term(&cdot, ncent)
137            .unwrap();
138        s.score_many(docs.iter().copied(), &mut out); // warm caches and branch predictors
139    }
140    for _ in 0..reps {
141        for arm in arms.iter_mut() {
142            let s = Scorer::new(&arm.lut, &q)
143                .with_centroid_term(&cdot, ncent)
144                .unwrap();
145            let t = Instant::now();
146            s.score_many(docs.iter().copied(), &mut out);
147            arm.ns.push(t.elapsed().as_nanos() as f64 / (ndocs * ntok) as f64);
148            arm.checksum = out.iter().map(|&v| v as f64).sum();
149        }
150    }
151
152    let reference = arms.last().expect("at least the scalar arm");
153    let (slow, want) = {
154        let mut v = reference.ns.clone();
155        v.sort_by(f64::total_cmp);
156        (v[0], reference.checksum)
157    };
158    let mut worst_spread = 0.0f64;
159    for arm in &arms {
160        let mut v = arm.ns.clone();
161        v.sort_by(f64::total_cmp);
162        let (best, med) = (v[0], median(&v));
163        let spread = (med - best) / best;
164        worst_spread = worst_spread.max(spread);
165        assert_eq!(
166            arm.checksum.to_bits(),
167            want.to_bits(),
168            "{}: checksum {} differs from the scalar reference {want}",
169            arm.label,
170            arm.checksum
171        );
172        println!(
173            "{:>26}: {best:7.2} ns/token  (median {med:7.2}, {:5.1} µs/doc, {:5.2}x scalar)",
174            arm.label,
175            best * ntok as f64 / 1e3,
176            slow / best,
177        );
178    }
179    println!(
180        "{:>26}: {:.1}% median-vs-best spread — {}",
181        "noise",
182        worst_spread * 100.0,
183        if worst_spread < 0.10 {
184            "quiet enough to compare kernels"
185        } else {
186            "TOO NOISY, differences under ~2x are not real; free the machine and rerun"
187        }
188    );
189}
Source

pub fn colbert(nbits: usize, bucket_weights: &[f32]) -> Result<Self, Error>

Convenience for the ColBERT / PLAID layout.

Examples found in repository?
examples/real_index.rs (line 170)
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}
Source

pub fn force_scalar(self, yes: bool) -> Self

Pin every score to the scalar reference kernel. For tests and for measuring what the SIMD is worth; the results are bit-identical either way. The environment variable MAXSIM_LUT_FORCE_SCALAR=1 has the same effect process-wide.

Examples found in repository?
examples/bench.rs (line 123)
55fn main() {
56    let args: Vec<usize> = std::env::args()
57        .skip(1)
58        .map(|a| a.parse().expect("integer arg"))
59        .collect();
60    let dim = args.first().copied().unwrap_or(128);
61    let nbits = args.get(1).copied().unwrap_or(4);
62    let nq = args.get(2).copied().unwrap_or(32);
63    let ntok = args.get(3).copied().unwrap_or(240);
64    let ndocs = args.get(4).copied().unwrap_or(1024);
65    let ncent = 16_384;
66    let reps = 9;
67    let mut rng = Rng(0x9E3779B97F4A7C15);
68
69    let p = ColbertPacking::new(nbits).unwrap();
70    let nb = 1usize << nbits;
71    let mut w: Vec<f32> = (0..nb).map(|_| rng.f32(-0.4, 0.4)).collect();
72    w.sort_by(|a, b| a.total_cmp(b));
73    let lut = Lut::new(&p, &w).unwrap();
74
75    let query: Vec<f32> = (0..nq * dim).map(|_| rng.f32(-1.0, 1.0)).collect();
76    let q = PreparedQuery::new(&lut, &query, nq, dim).unwrap();
77    let cdot: Vec<f32> = (0..ncent * nq).map(|_| rng.f32(-1.0, 1.0)).collect();
78
79    let pdim = dim / p.keys_per_byte();
80    let mut packed = vec![0u8; ndocs * ntok * pdim];
81    for b in packed.iter_mut() {
82        *b = (rng.next() >> 56) as u8;
83    }
84    let codes: Vec<u32> = (0..ndocs * ntok)
85        .map(|_| (rng.next() % ncent as u64) as u32)
86        .collect();
87    let inv: Vec<f32> = (0..ndocs * ntok).map(|_| rng.f32(0.8, 1.2)).collect();
88    let docs: Vec<DocView> = (0..ndocs)
89        .map(|d| {
90            DocView::new(&packed[d * ntok * pdim..(d + 1) * ntok * pdim], ntok, pdim)
91                .codes(Codes::U32(&codes[d * ntok..(d + 1) * ntok]))
92                .inv_norms(&inv[d * ntok..(d + 1) * ntok])
93        })
94        .collect();
95
96    // One arm per executable kernel, plus the scalar reference. `dispatch`
97    // is what an unpinned host would get, named so the calibrated choice is
98    // visible next to the kernels it chose between.
99    let dispatched = lut.kernel(dim);
100    let mut arms: Vec<Arm> = Vec::new();
101    for &k in supported_kernels() {
102        arms.push(Arm {
103            label: if k == dispatched {
104                format!("{k} (dispatched)")
105            } else {
106                format!("{k}")
107            },
108            lut: lut.clone().pin_kernel(Some(k)),
109            ns: Vec::new(),
110            checksum: 0.0,
111        });
112    }
113    if !dispatched.is_simd() {
114        arms.push(Arm {
115            label: format!("{dispatched} (dispatched)"),
116            lut: lut.clone(),
117            ns: Vec::new(),
118            checksum: 0.0,
119        });
120    }
121    arms.push(Arm {
122        label: "scalar reference".to_string(),
123        lut: lut.clone().force_scalar(true),
124        ns: Vec::new(),
125        checksum: 0.0,
126    });
127
128    println!(
129        "dim {dim}, nbits {nbits}, {nq} query tokens, {ndocs} docs × {ntok} tokens, {ncent} centroids\narch {}, dispatched kernel: {dispatched}, {reps} interleaved rounds",
130        std::env::consts::ARCH,
131    );
132
133    let mut out = vec![0.0f32; ndocs];
134    for arm in arms.iter_mut() {
135        let s = Scorer::new(&arm.lut, &q)
136            .with_centroid_term(&cdot, ncent)
137            .unwrap();
138        s.score_many(docs.iter().copied(), &mut out); // warm caches and branch predictors
139    }
140    for _ in 0..reps {
141        for arm in arms.iter_mut() {
142            let s = Scorer::new(&arm.lut, &q)
143                .with_centroid_term(&cdot, ncent)
144                .unwrap();
145            let t = Instant::now();
146            s.score_many(docs.iter().copied(), &mut out);
147            arm.ns.push(t.elapsed().as_nanos() as f64 / (ndocs * ntok) as f64);
148            arm.checksum = out.iter().map(|&v| v as f64).sum();
149        }
150    }
151
152    let reference = arms.last().expect("at least the scalar arm");
153    let (slow, want) = {
154        let mut v = reference.ns.clone();
155        v.sort_by(f64::total_cmp);
156        (v[0], reference.checksum)
157    };
158    let mut worst_spread = 0.0f64;
159    for arm in &arms {
160        let mut v = arm.ns.clone();
161        v.sort_by(f64::total_cmp);
162        let (best, med) = (v[0], median(&v));
163        let spread = (med - best) / best;
164        worst_spread = worst_spread.max(spread);
165        assert_eq!(
166            arm.checksum.to_bits(),
167            want.to_bits(),
168            "{}: checksum {} differs from the scalar reference {want}",
169            arm.label,
170            arm.checksum
171        );
172        println!(
173            "{:>26}: {best:7.2} ns/token  (median {med:7.2}, {:5.1} µs/doc, {:5.2}x scalar)",
174            arm.label,
175            best * ntok as f64 / 1e3,
176            slow / best,
177        );
178    }
179    println!(
180        "{:>26}: {:.1}% median-vs-best spread — {}",
181        "noise",
182        worst_spread * 100.0,
183        if worst_spread < 0.10 {
184            "quiet enough to compare kernels"
185        } else {
186            "TOO NOISY, differences under ~2x are not real; free the machine and rerun"
187        }
188    );
189}
Source

pub fn pin_kernel(self, kernel: Option<Kernel>) -> Self

Pin dispatch to one kernel instead of the calibrated choice.

Only a kernel this CPU can execute is honoured (check with crate::supported_kernels); anything else is ignored and dispatch proceeds normally, because a kernel the CPU cannot run has no meaningful behaviour to fall back to. Scoring is bit-identical whichever kernel runs, so this only affects speed.

Use it to benchmark one path, or to make dispatch deterministic on a fleet of mixed cores. None restores the default.

Examples found in repository?
examples/bench.rs (line 108)
55fn main() {
56    let args: Vec<usize> = std::env::args()
57        .skip(1)
58        .map(|a| a.parse().expect("integer arg"))
59        .collect();
60    let dim = args.first().copied().unwrap_or(128);
61    let nbits = args.get(1).copied().unwrap_or(4);
62    let nq = args.get(2).copied().unwrap_or(32);
63    let ntok = args.get(3).copied().unwrap_or(240);
64    let ndocs = args.get(4).copied().unwrap_or(1024);
65    let ncent = 16_384;
66    let reps = 9;
67    let mut rng = Rng(0x9E3779B97F4A7C15);
68
69    let p = ColbertPacking::new(nbits).unwrap();
70    let nb = 1usize << nbits;
71    let mut w: Vec<f32> = (0..nb).map(|_| rng.f32(-0.4, 0.4)).collect();
72    w.sort_by(|a, b| a.total_cmp(b));
73    let lut = Lut::new(&p, &w).unwrap();
74
75    let query: Vec<f32> = (0..nq * dim).map(|_| rng.f32(-1.0, 1.0)).collect();
76    let q = PreparedQuery::new(&lut, &query, nq, dim).unwrap();
77    let cdot: Vec<f32> = (0..ncent * nq).map(|_| rng.f32(-1.0, 1.0)).collect();
78
79    let pdim = dim / p.keys_per_byte();
80    let mut packed = vec![0u8; ndocs * ntok * pdim];
81    for b in packed.iter_mut() {
82        *b = (rng.next() >> 56) as u8;
83    }
84    let codes: Vec<u32> = (0..ndocs * ntok)
85        .map(|_| (rng.next() % ncent as u64) as u32)
86        .collect();
87    let inv: Vec<f32> = (0..ndocs * ntok).map(|_| rng.f32(0.8, 1.2)).collect();
88    let docs: Vec<DocView> = (0..ndocs)
89        .map(|d| {
90            DocView::new(&packed[d * ntok * pdim..(d + 1) * ntok * pdim], ntok, pdim)
91                .codes(Codes::U32(&codes[d * ntok..(d + 1) * ntok]))
92                .inv_norms(&inv[d * ntok..(d + 1) * ntok])
93        })
94        .collect();
95
96    // One arm per executable kernel, plus the scalar reference. `dispatch`
97    // is what an unpinned host would get, named so the calibrated choice is
98    // visible next to the kernels it chose between.
99    let dispatched = lut.kernel(dim);
100    let mut arms: Vec<Arm> = Vec::new();
101    for &k in supported_kernels() {
102        arms.push(Arm {
103            label: if k == dispatched {
104                format!("{k} (dispatched)")
105            } else {
106                format!("{k}")
107            },
108            lut: lut.clone().pin_kernel(Some(k)),
109            ns: Vec::new(),
110            checksum: 0.0,
111        });
112    }
113    if !dispatched.is_simd() {
114        arms.push(Arm {
115            label: format!("{dispatched} (dispatched)"),
116            lut: lut.clone(),
117            ns: Vec::new(),
118            checksum: 0.0,
119        });
120    }
121    arms.push(Arm {
122        label: "scalar reference".to_string(),
123        lut: lut.clone().force_scalar(true),
124        ns: Vec::new(),
125        checksum: 0.0,
126    });
127
128    println!(
129        "dim {dim}, nbits {nbits}, {nq} query tokens, {ndocs} docs × {ntok} tokens, {ncent} centroids\narch {}, dispatched kernel: {dispatched}, {reps} interleaved rounds",
130        std::env::consts::ARCH,
131    );
132
133    let mut out = vec![0.0f32; ndocs];
134    for arm in arms.iter_mut() {
135        let s = Scorer::new(&arm.lut, &q)
136            .with_centroid_term(&cdot, ncent)
137            .unwrap();
138        s.score_many(docs.iter().copied(), &mut out); // warm caches and branch predictors
139    }
140    for _ in 0..reps {
141        for arm in arms.iter_mut() {
142            let s = Scorer::new(&arm.lut, &q)
143                .with_centroid_term(&cdot, ncent)
144                .unwrap();
145            let t = Instant::now();
146            s.score_many(docs.iter().copied(), &mut out);
147            arm.ns.push(t.elapsed().as_nanos() as f64 / (ndocs * ntok) as f64);
148            arm.checksum = out.iter().map(|&v| v as f64).sum();
149        }
150    }
151
152    let reference = arms.last().expect("at least the scalar arm");
153    let (slow, want) = {
154        let mut v = reference.ns.clone();
155        v.sort_by(f64::total_cmp);
156        (v[0], reference.checksum)
157    };
158    let mut worst_spread = 0.0f64;
159    for arm in &arms {
160        let mut v = arm.ns.clone();
161        v.sort_by(f64::total_cmp);
162        let (best, med) = (v[0], median(&v));
163        let spread = (med - best) / best;
164        worst_spread = worst_spread.max(spread);
165        assert_eq!(
166            arm.checksum.to_bits(),
167            want.to_bits(),
168            "{}: checksum {} differs from the scalar reference {want}",
169            arm.label,
170            arm.checksum
171        );
172        println!(
173            "{:>26}: {best:7.2} ns/token  (median {med:7.2}, {:5.1} µs/doc, {:5.2}x scalar)",
174            arm.label,
175            best * ntok as f64 / 1e3,
176            slow / best,
177        );
178    }
179    println!(
180        "{:>26}: {:.1}% median-vs-best spread — {}",
181        "noise",
182        worst_spread * 100.0,
183        if worst_spread < 0.10 {
184            "quiet enough to compare kernels"
185        } else {
186            "TOO NOISY, differences under ~2x are not real; free the machine and rerun"
187        }
188    );
189}
Source

pub fn nbits(&self) -> usize

Code width this table was built for.

Source

pub fn keys_per_byte(&self) -> usize

8 / nbits: how many dims one packed byte carries.

Source

pub fn scale(&self) -> f32

Dequantisation scale: fused as f32 · scale ≈ bucket_weight.

Source

pub fn expand(&self, byte: u8) -> &[i8]

The 8/nbits int8 weights a packed byte expands to, in dim order.

Source

pub fn fused_table(&self) -> &[i8]

The whole fused table, [256 · keys_per_byte], row b = byte b.

Source

pub fn nibble_tables(&self) -> Option<&NibbleTables>

The nibble-factored tables, if the layout admits them (always, for crate::ColbertPacking at nbits 1, 2 or 4; never at nbits 8).

Source

pub fn kernel(&self, dim: usize) -> Kernel

Which kernel crate::Scorer::score will run for this table and dim on this CPU. Print it next to any benchmark number: a speedup attributed to a path that never executed is the easiest measurement error to make and the hardest to notice.

Examples found in repository?
examples/bench.rs (line 99)
55fn main() {
56    let args: Vec<usize> = std::env::args()
57        .skip(1)
58        .map(|a| a.parse().expect("integer arg"))
59        .collect();
60    let dim = args.first().copied().unwrap_or(128);
61    let nbits = args.get(1).copied().unwrap_or(4);
62    let nq = args.get(2).copied().unwrap_or(32);
63    let ntok = args.get(3).copied().unwrap_or(240);
64    let ndocs = args.get(4).copied().unwrap_or(1024);
65    let ncent = 16_384;
66    let reps = 9;
67    let mut rng = Rng(0x9E3779B97F4A7C15);
68
69    let p = ColbertPacking::new(nbits).unwrap();
70    let nb = 1usize << nbits;
71    let mut w: Vec<f32> = (0..nb).map(|_| rng.f32(-0.4, 0.4)).collect();
72    w.sort_by(|a, b| a.total_cmp(b));
73    let lut = Lut::new(&p, &w).unwrap();
74
75    let query: Vec<f32> = (0..nq * dim).map(|_| rng.f32(-1.0, 1.0)).collect();
76    let q = PreparedQuery::new(&lut, &query, nq, dim).unwrap();
77    let cdot: Vec<f32> = (0..ncent * nq).map(|_| rng.f32(-1.0, 1.0)).collect();
78
79    let pdim = dim / p.keys_per_byte();
80    let mut packed = vec![0u8; ndocs * ntok * pdim];
81    for b in packed.iter_mut() {
82        *b = (rng.next() >> 56) as u8;
83    }
84    let codes: Vec<u32> = (0..ndocs * ntok)
85        .map(|_| (rng.next() % ncent as u64) as u32)
86        .collect();
87    let inv: Vec<f32> = (0..ndocs * ntok).map(|_| rng.f32(0.8, 1.2)).collect();
88    let docs: Vec<DocView> = (0..ndocs)
89        .map(|d| {
90            DocView::new(&packed[d * ntok * pdim..(d + 1) * ntok * pdim], ntok, pdim)
91                .codes(Codes::U32(&codes[d * ntok..(d + 1) * ntok]))
92                .inv_norms(&inv[d * ntok..(d + 1) * ntok])
93        })
94        .collect();
95
96    // One arm per executable kernel, plus the scalar reference. `dispatch`
97    // is what an unpinned host would get, named so the calibrated choice is
98    // visible next to the kernels it chose between.
99    let dispatched = lut.kernel(dim);
100    let mut arms: Vec<Arm> = Vec::new();
101    for &k in supported_kernels() {
102        arms.push(Arm {
103            label: if k == dispatched {
104                format!("{k} (dispatched)")
105            } else {
106                format!("{k}")
107            },
108            lut: lut.clone().pin_kernel(Some(k)),
109            ns: Vec::new(),
110            checksum: 0.0,
111        });
112    }
113    if !dispatched.is_simd() {
114        arms.push(Arm {
115            label: format!("{dispatched} (dispatched)"),
116            lut: lut.clone(),
117            ns: Vec::new(),
118            checksum: 0.0,
119        });
120    }
121    arms.push(Arm {
122        label: "scalar reference".to_string(),
123        lut: lut.clone().force_scalar(true),
124        ns: Vec::new(),
125        checksum: 0.0,
126    });
127
128    println!(
129        "dim {dim}, nbits {nbits}, {nq} query tokens, {ndocs} docs × {ntok} tokens, {ncent} centroids\narch {}, dispatched kernel: {dispatched}, {reps} interleaved rounds",
130        std::env::consts::ARCH,
131    );
132
133    let mut out = vec![0.0f32; ndocs];
134    for arm in arms.iter_mut() {
135        let s = Scorer::new(&arm.lut, &q)
136            .with_centroid_term(&cdot, ncent)
137            .unwrap();
138        s.score_many(docs.iter().copied(), &mut out); // warm caches and branch predictors
139    }
140    for _ in 0..reps {
141        for arm in arms.iter_mut() {
142            let s = Scorer::new(&arm.lut, &q)
143                .with_centroid_term(&cdot, ncent)
144                .unwrap();
145            let t = Instant::now();
146            s.score_many(docs.iter().copied(), &mut out);
147            arm.ns.push(t.elapsed().as_nanos() as f64 / (ndocs * ntok) as f64);
148            arm.checksum = out.iter().map(|&v| v as f64).sum();
149        }
150    }
151
152    let reference = arms.last().expect("at least the scalar arm");
153    let (slow, want) = {
154        let mut v = reference.ns.clone();
155        v.sort_by(f64::total_cmp);
156        (v[0], reference.checksum)
157    };
158    let mut worst_spread = 0.0f64;
159    for arm in &arms {
160        let mut v = arm.ns.clone();
161        v.sort_by(f64::total_cmp);
162        let (best, med) = (v[0], median(&v));
163        let spread = (med - best) / best;
164        worst_spread = worst_spread.max(spread);
165        assert_eq!(
166            arm.checksum.to_bits(),
167            want.to_bits(),
168            "{}: checksum {} differs from the scalar reference {want}",
169            arm.label,
170            arm.checksum
171        );
172        println!(
173            "{:>26}: {best:7.2} ns/token  (median {med:7.2}, {:5.1} µs/doc, {:5.2}x scalar)",
174            arm.label,
175            best * ntok as f64 / 1e3,
176            slow / best,
177        );
178    }
179    println!(
180        "{:>26}: {:.1}% median-vs-best spread — {}",
181        "noise",
182        worst_spread * 100.0,
183        if worst_spread < 0.10 {
184            "quiet enough to compare kernels"
185        } else {
186            "TOO NOISY, differences under ~2x are not real; free the machine and rerun"
187        }
188    );
189}
More examples
Hide additional examples
examples/real_index.rs (line 176)
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}

Trait Implementations§

Source§

impl Clone for Lut

Source§

fn clone(&self) -> Lut

Returns a duplicate of the value. Read more
1.0.0 (const: unstable) · Source§

fn clone_from(&mut self, source: &Self)

Performs copy-assignment from source. Read more
Source§

impl Debug for Lut

Source§

fn fmt(&self, f: &mut Formatter<'_>) -> Result

Formats the value using the given formatter. Read more

Auto Trait Implementations§

§

impl Freeze for Lut

§

impl RefUnwindSafe for Lut

§

impl Send for Lut

§

impl Sync for Lut

§

impl Unpin for Lut

§

impl UnsafeUnpin for Lut

§

impl UnwindSafe for Lut

Blanket Implementations§

Source§

impl<T> Any for T
where T: 'static + ?Sized,

Source§

fn type_id(&self) -> TypeId

Gets the TypeId of self. Read more
Source§

impl<T> Borrow<T> for T
where T: ?Sized,

Source§

fn borrow(&self) -> &T

Immutably borrows from an owned value. Read more
Source§

impl<T> BorrowMut<T> for T
where T: ?Sized,

Source§

fn borrow_mut(&mut self) -> &mut T

Mutably borrows from an owned value. Read more
Source§

impl<T> CloneToUninit for T
where T: Clone,

Source§

unsafe fn clone_to_uninit(&self, dest: *mut u8)

🔬This is a nightly-only experimental API. (clone_to_uninit)
Performs copy-assignment from self to dest. Read more
Source§

impl<T> From<T> for T

Source§

fn from(t: T) -> T

Returns the argument unchanged.

Source§

impl<T, U> Into<U> for T
where U: From<T>,

Source§

fn into(self) -> U

Calls U::from(self).

That is, this conversion is whatever the implementation of From<T> for U chooses to do.

Source§

impl<T> ToOwned for T
where T: Clone,

Source§

type Owned = T

The resulting type after obtaining ownership.
Source§

fn to_owned(&self) -> T

Creates owned data from borrowed data, usually by cloning. Read more
Source§

fn clone_into(&self, target: &mut T)

Uses borrowed data to replace owned data, usually by cloning. Read more
Source§

impl<T, U> TryFrom<U> for T
where U: Into<T>,

Source§

type Error = !

The type returned in the event of a conversion error.
Source§

fn try_from(value: U) -> Result<T, !>

Performs the conversion.
Source§

impl<T, U> TryInto<U> for T
where U: TryFrom<T>,

Source§

type Error = <U as TryFrom<T>>::Error

The type returned in the event of a conversion error.
Source§

fn try_into(self) -> Result<U, <U as TryFrom<T>>::Error>

Performs the conversion.