1use crate::{
18 Result,
19 codec::{EncodedVector, ResidualCodec, train_quantizer},
20 kmeans::{assign_points, fit},
21};
22
23#[derive(Debug, Clone, Default)]
37pub struct InvertedFile {
38 pub lists: Vec<Vec<u32>>,
41}
42
43impl InvertedFile {
44 pub fn num_centroids(&self) -> usize {
46 self.lists.len()
47 }
48
49 pub fn docs_for_centroid(&self, centroid_id: usize) -> &[u32] {
53 self.lists
54 .get(centroid_id)
55 .map(Vec::as_slice)
56 .unwrap_or(&[])
57 }
58
59 pub fn total_doc_postings(&self) -> usize {
65 self.lists.iter().map(Vec::len).sum()
66 }
67}
68
69#[derive(Debug, Clone)]
75pub struct DocumentTokens {
76 pub doc_id: u64,
77 pub tokens: Vec<f32>,
78 pub n_tokens: usize,
79}
80
81impl DocumentTokens {
82 pub fn flat_len(&self) -> usize {
84 self.tokens.len()
85 }
86}
87
88#[derive(Debug, Clone, Copy)]
90pub struct IndexParams {
91 pub dim: usize,
93 pub nbits: u32,
95 pub k_centroids: usize,
97 pub max_kmeans_iters: usize,
99}
100
101#[derive(Debug, Clone)]
108pub struct Index {
109 pub params: IndexParams,
110 pub codec: ResidualCodec,
111 pub doc_ids: Vec<u64>,
112 pub doc_tokens: Vec<Vec<EncodedVector>>,
114 pub ivf: InvertedFile,
116}
117
118impl Index {
119 pub fn num_documents(&self) -> usize {
121 self.doc_ids.len()
122 }
123
124 pub fn num_tokens(&self) -> usize {
126 self.doc_tokens.iter().map(Vec::len).sum()
127 }
128
129 pub fn position_of(&self, doc_id: u64) -> Option<usize> {
131 self.doc_ids.iter().position(|id| *id == doc_id)
132 }
133}
134
135pub fn build_index(
156 documents: &[DocumentTokens],
157 params: IndexParams,
158) -> Result<Index> {
159 assert!(params.dim > 0, "build_index: dim must be positive");
160 assert!(
161 params.k_centroids > 0,
162 "build_index: k_centroids must be positive"
163 );
164 assert!(
165 params.nbits > 0 && params.nbits <= 8,
166 "build_index: nbits must be in 1..=8, got {}",
167 params.nbits,
168 );
169
170 for doc in documents {
171 assert!(
172 doc.tokens.len() == doc.n_tokens * params.dim,
173 "build_index: doc {} declared {} tokens but carries {} f32s (dim={})",
174 doc.doc_id,
175 doc.n_tokens,
176 doc.tokens.len(),
177 params.dim,
178 );
179 }
180
181 let total_tokens: usize = documents.iter().map(|d| d.n_tokens).sum();
183 assert!(
184 total_tokens >= params.k_centroids,
185 "build_index: need at least {} tokens for {} centroids, got {}",
186 params.k_centroids,
187 params.k_centroids,
188 total_tokens,
189 );
190 let mut pool: Vec<f32> = Vec::with_capacity(total_tokens * params.dim);
191 for doc in documents {
192 pool.extend_from_slice(&doc.tokens);
193 }
194
195 let centroids = fit(
197 &pool,
198 params.k_centroids,
199 params.dim,
200 params.max_kmeans_iters.max(1),
201 )?;
202
203 let assignments = assign_points(&pool, ¢roids, params.dim)?;
206 let mut residual_sample: Vec<f32> = Vec::with_capacity(pool.len());
207 for (token, &cluster) in pool.chunks_exact(params.dim).zip(&assignments) {
208 let centroid =
209 ¢roids[cluster * params.dim..(cluster + 1) * params.dim];
210 for (t, c) in token.iter().zip(centroid) {
211 residual_sample.push(*t - *c);
212 }
213 }
214
215 let (bucket_cutoffs, bucket_weights) =
217 train_quantizer(&residual_sample, params.nbits);
218
219 let codec = ResidualCodec {
220 nbits: params.nbits,
221 dim: params.dim,
222 centroids,
223 bucket_cutoffs,
224 bucket_weights,
225 };
226 codec.validate()?;
227
228 let (all_centroid_ids, all_codes) = codec.batch_encode_tokens(&pool)?;
233 let packed_per_token = codec.packed_bytes();
234
235 let mut doc_ids = Vec::with_capacity(documents.len());
236 let mut doc_tokens = Vec::with_capacity(documents.len());
237 let mut token_offset = 0usize;
238 for doc in documents {
239 doc_ids.push(doc.doc_id);
240 let n_tok = doc.n_tokens;
241 let cids = &all_centroid_ids[token_offset..token_offset + n_tok];
242 let codes_slice = &all_codes[token_offset * packed_per_token
243 ..(token_offset + n_tok) * packed_per_token];
244 let encoded: Vec<EncodedVector> = (0..n_tok)
245 .map(|i| EncodedVector {
246 centroid_id: cids[i],
247 codes: codes_slice
248 [i * packed_per_token..(i + 1) * packed_per_token]
249 .to_vec(),
250 })
251 .collect();
252 doc_tokens.push(encoded);
253 token_offset += n_tok;
254 }
255
256 let ivf = build_inverted_file(&doc_tokens, params.k_centroids);
257
258 Ok(Index {
259 params,
260 codec,
261 doc_ids,
262 doc_tokens,
263 ivf,
264 })
265}
266
267pub(crate) fn build_inverted_file(
271 doc_tokens: &[Vec<EncodedVector>],
272 k_centroids: usize,
273) -> InvertedFile {
274 let mut lists: Vec<Vec<u32>> = vec![Vec::new(); k_centroids];
275 for (doc_idx, encoded) in doc_tokens.iter().enumerate() {
276 let mut touched: Vec<u32> =
277 encoded.iter().map(|ev| ev.centroid_id).collect();
278 touched.sort_unstable();
279 touched.dedup();
280 for cid in touched {
281 lists[cid as usize].push(doc_idx as u32);
282 }
283 }
284 InvertedFile { lists }
285}
286
287#[cfg(test)]
288mod tests {
289 use super::*;
290 use crate::distance::squared_l2;
291
292 fn small_corpus() -> Vec<DocumentTokens> {
294 vec![
295 DocumentTokens {
296 doc_id: 1,
297 tokens: vec![0.0, 0.0, 0.1, 0.2, -0.1, 0.1],
298 n_tokens: 3,
299 },
300 DocumentTokens {
301 doc_id: 2,
302 tokens: vec![10.0, 10.0, 10.2, 9.9, 9.8, 10.1],
303 n_tokens: 3,
304 },
305 DocumentTokens {
306 doc_id: 3,
307 tokens: vec![0.3, -0.2, 9.7, 10.2],
308 n_tokens: 2,
309 },
310 ]
311 }
312
313 fn default_params() -> IndexParams {
314 IndexParams {
315 dim: 2,
316 nbits: 2,
317 k_centroids: 2,
318 max_kmeans_iters: 50,
319 }
320 }
321
322 #[test]
323 fn build_index_encodes_every_token() {
324 let docs = small_corpus();
325 let params = default_params();
326 let expected_total: usize = docs.iter().map(|d| d.n_tokens).sum();
327
328 let index = build_index(&docs, params).unwrap();
329
330 assert_eq!(index.num_documents(), docs.len());
331 assert_eq!(index.num_tokens(), expected_total);
332 for (encoded, doc) in index.doc_tokens.iter().zip(docs.iter()) {
333 assert_eq!(encoded.len(), doc.n_tokens);
334 }
335 }
336
337 #[test]
338 fn build_index_assigns_tokens_in_each_cluster_to_the_closest_centroid() {
339 let docs = small_corpus();
340 let params = default_params();
341 let index = build_index(&docs, params).unwrap();
342
343 let c0 = &index.codec.centroids[0..2];
346 let c1 = &index.codec.centroids[2..4];
347
348 let (near_origin, near_ten) =
349 if squared_l2(c0, &[0.0, 0.0]) < squared_l2(c0, &[10.0, 10.0]) {
350 (c0, c1)
351 } else {
352 (c1, c0)
353 };
354
355 assert!(squared_l2(near_origin, &[0.0, 0.0]) < 1.0);
356 assert!(squared_l2(near_ten, &[10.0, 10.0]) < 1.0);
357 }
358
359 #[test]
360 fn build_index_round_trip_reconstruction_error_is_bounded() {
361 let docs = small_corpus();
362 let params = default_params();
363 let index = build_index(&docs, params).unwrap();
364
365 for (doc, encoded_doc) in docs.iter().zip(index.doc_tokens.iter()) {
366 for (token, encoded) in
367 doc.tokens.chunks_exact(params.dim).zip(encoded_doc.iter())
368 {
369 let decoded = index.codec.decode_vector(encoded).unwrap();
370 let err = squared_l2(token, &decoded).sqrt();
371 assert!(
374 err < 0.6,
375 "reconstruction error {err} too large for token {token:?}"
376 );
377 }
378 }
379 }
380
381 #[test]
382 fn build_index_preserves_document_id_order() {
383 let docs = small_corpus();
384 let index = build_index(&docs, default_params()).unwrap();
385 let expected_ids: Vec<u64> = docs.iter().map(|d| d.doc_id).collect();
386 assert_eq!(index.doc_ids, expected_ids);
387 assert_eq!(index.position_of(2), Some(1));
388 assert_eq!(index.position_of(999), None);
389 }
390
391 #[test]
392 fn build_index_handles_document_with_no_tokens() {
393 let mut docs = small_corpus();
396 docs.push(DocumentTokens {
397 doc_id: 42,
398 tokens: vec![],
399 n_tokens: 0,
400 });
401 let index = build_index(&docs, default_params()).unwrap();
402 assert_eq!(index.num_documents(), 4);
403 assert_eq!(index.doc_tokens[3].len(), 0);
404 }
405
406 #[test]
407 #[should_panic(expected = "declared")]
408 fn build_index_panics_on_mismatched_token_count() {
409 let docs = vec![DocumentTokens {
410 doc_id: 1,
411 tokens: vec![0.0, 0.0, 1.0],
412 n_tokens: 2, }];
414 let _ = build_index(&docs, default_params()).unwrap();
415 }
416
417 #[test]
418 fn build_index_ivf_has_one_list_per_centroid() {
419 let docs = small_corpus();
420 let index = build_index(&docs, default_params()).unwrap();
421
422 assert_eq!(
423 index.ivf.num_centroids(),
424 default_params().k_centroids,
425 "IVF has one list per centroid",
426 );
427 }
428
429 #[test]
430 fn build_index_ivf_postings_cover_every_doc_that_touches_each_centroid() {
431 let docs = small_corpus();
435 let index = build_index(&docs, default_params()).unwrap();
436
437 for (doc_idx, encoded_doc) in index.doc_tokens.iter().enumerate() {
438 for ev in encoded_doc {
439 let postings =
440 index.ivf.docs_for_centroid(ev.centroid_id as usize);
441 assert!(
442 postings.contains(&(doc_idx as u32)),
443 "doc_idx={doc_idx} missing from centroid {} postings",
444 ev.centroid_id,
445 );
446 }
447 }
448 }
449
450 #[test]
451 fn build_index_ivf_postings_are_unique_per_centroid() {
452 let docs = small_corpus();
456 let index = build_index(&docs, default_params()).unwrap();
457
458 for c in 0..index.ivf.num_centroids() {
459 let postings = index.ivf.docs_for_centroid(c);
460 let mut unique: Vec<u32> = postings.to_vec();
461 unique.sort_unstable();
462 unique.dedup();
463 assert_eq!(
464 unique.len(),
465 postings.len(),
466 "centroid {c} has duplicate doc entries: {postings:?}",
467 );
468 }
469 }
470
471 #[test]
472 fn build_index_dedupes_repeated_tokens_in_same_centroid() {
473 let docs = vec![
477 DocumentTokens {
478 doc_id: 1,
479 tokens: vec![0.0, 0.0, 0.05, -0.02, -0.03, 0.01],
480 n_tokens: 3,
481 },
482 DocumentTokens {
483 doc_id: 2,
484 tokens: vec![10.0, 10.0, 10.1, 9.9],
485 n_tokens: 2,
486 },
487 ];
488 let index = build_index(&docs, default_params()).unwrap();
489
490 for c in 0..index.ivf.num_centroids() {
491 let postings = index.ivf.docs_for_centroid(c);
492 let count_of_doc_0 = postings.iter().filter(|&&d| d == 0).count();
493 assert!(
494 count_of_doc_0 <= 1,
495 "doc 0 appears {count_of_doc_0} times in centroid {c}",
496 );
497 }
498 }
499
500 #[test]
501 fn inverted_file_out_of_range_returns_empty_slice() {
502 let ivf = InvertedFile {
503 lists: vec![vec![0u32]],
504 };
505 assert_eq!(ivf.docs_for_centroid(0).len(), 1);
506 assert!(ivf.docs_for_centroid(999).is_empty());
507 }
508
509 #[test]
510 #[should_panic(expected = "at least")]
511 fn build_index_panics_when_too_few_tokens_for_k() {
512 let docs = vec![DocumentTokens {
513 doc_id: 1,
514 tokens: vec![0.0, 1.0],
515 n_tokens: 1,
516 }];
517 let params = IndexParams {
518 dim: 2,
519 nbits: 2,
520 k_centroids: 4,
521 max_kmeans_iters: 10,
522 };
523 let _ = build_index(&docs, params).unwrap();
524 }
525}