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
impl Lut
Sourcepub fn new<P: Packing>(
packing: &P,
bucket_weights: &[f32],
) -> Result<Self, Error>
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?
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}Sourcepub fn colbert(nbits: usize, bucket_weights: &[f32]) -> Result<Self, Error>
pub fn colbert(nbits: usize, bucket_weights: &[f32]) -> Result<Self, Error>
Convenience for the ColBERT / PLAID layout.
Examples found in repository?
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 = ¢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 // 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 = ¢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 // 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 = ¢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}Sourcepub fn force_scalar(self, yes: bool) -> Self
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?
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}Sourcepub fn pin_kernel(self, kernel: Option<Kernel>) -> Self
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?
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}Sourcepub fn keys_per_byte(&self) -> usize
pub fn keys_per_byte(&self) -> usize
8 / nbits: how many dims one packed byte carries.
Sourcepub fn expand(&self, byte: u8) -> &[i8]
pub fn expand(&self, byte: u8) -> &[i8]
The 8/nbits int8 weights a packed byte expands to, in dim order.
Sourcepub fn fused_table(&self) -> &[i8]
pub fn fused_table(&self) -> &[i8]
The whole fused table, [256 · keys_per_byte], row b = byte b.
Sourcepub fn nibble_tables(&self) -> Option<&NibbleTables>
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).
Sourcepub fn kernel(&self, dim: usize) -> Kernel
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?
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
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 = ¢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 // 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 = ¢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 // 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 = ¢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}