docbert_plaid/index.rs
1//! Index construction: turn a corpus of token embeddings into a searchable
2//! PLAID index.
3//!
4//! `build_index` ties together the three lower layers:
5//!
6//! 1. Flatten every document's token matrix into one big cloud of points.
7//! 2. Run k-means to pick coarse centroids ([`crate::kmeans::fit`]).
8//! 3. Assign every token to a centroid, compute residuals, train cutoffs
9//! and weights on the residuals ([`crate::codec::train_quantizer`]),
10//! and encode each token against the fresh codec.
11//!
12//! The resulting [`Index`] keeps one [`EncodedVector`] per original token
13//! plus the per-document `doc_id` bookkeeping. Inverted-file construction
14//! and the query-time search path will be added on top in later
15//! TDD cycles.
16
17use crate::{
18 Result,
19 codec::{EncodedVector, ResidualCodec, train_quantizer},
20 kmeans::{assign_points, fit},
21};
22
23/// Cap on tokens used to train the residual quantizer.
24///
25/// Quantile estimation for `2^nbits` bucket cutoffs converges well
26/// below 100k samples; training on the full corpus (potentially
27/// millions of tokens × the embedding dim, so gigabytes of residuals)
28/// adds no statistical benefit while pushing peak RSS past what a
29/// host can absorb on large collections. 65,536 tokens × 128 dims is
30/// ~8M residual values, which still over-samples every cutoff by
31/// several orders of magnitude and sorts in under a second.
32const MAX_QUANTIZER_TRAINING_TOKENS: usize = 65_536;
33
34/// Inverted file: for each centroid, the sorted list of unique document
35/// indices that have at least one token clustered in that centroid.
36///
37/// This mirrors PLAID's "centroid → unique passage ids" layout from
38/// §3 of the paper: candidate generation only needs to know which
39/// documents are reachable via a probed centroid, and deduplicating
40/// per-doc keeps the posting lists small even when a single document
41/// has many tokens mapped to the same cluster.
42///
43/// The search path uses this to expand a query token to a shortlist of
44/// document candidates: find the centroids with the highest dot-product
45/// against the query token, then gather every document listed under
46/// those centroids.
47#[derive(Debug, Clone, Default)]
48pub struct InvertedFile {
49 /// `lists[c]` holds the sorted, deduplicated `doc_idx`s of every
50 /// document with at least one token assigned to centroid `c`.
51 pub lists: Vec<Vec<u32>>,
52}
53
54impl InvertedFile {
55 /// Total number of centroids the IVF spans.
56 pub fn num_centroids(&self) -> usize {
57 self.lists.len()
58 }
59
60 /// Document indices currently associated with `centroid_id`, or an
61 /// empty slice if the centroid is out of range. Entries are sorted
62 /// ascending and contain no duplicates.
63 pub fn docs_for_centroid(&self, centroid_id: usize) -> &[u32] {
64 self.lists
65 .get(centroid_id)
66 .map(Vec::as_slice)
67 .unwrap_or(&[])
68 }
69
70 /// Total number of (centroid, doc) postings across every list.
71 ///
72 /// This is the sum of `lists[c].len()` over all centroids `c`. It
73 /// is at most `num_centroids * num_documents` and at least equal
74 /// to the number of documents that contain any tokens at all.
75 pub fn total_doc_postings(&self) -> usize {
76 self.lists.iter().map(Vec::len).sum()
77 }
78}
79
80/// A single document's worth of token embeddings, ready to index.
81///
82/// `tokens` is a flat row-major `n_tokens × dim` buffer. Keeping the
83/// tokens flat mirrors the way docbert already stores ColBERT outputs in
84/// `embeddings.db` and avoids an intermediate `Vec<Vec<f32>>` allocation.
85#[derive(Debug, Clone)]
86pub struct DocumentTokens {
87 pub doc_id: u64,
88 pub tokens: Vec<f32>,
89 pub n_tokens: usize,
90}
91
92impl DocumentTokens {
93 /// Total number of f32 values this document contributes.
94 pub fn flat_len(&self) -> usize {
95 self.tokens.len()
96 }
97}
98
99/// Parameters that control how an [`Index`] is built.
100#[derive(Debug, Clone, Copy)]
101pub struct IndexParams {
102 /// Dimensionality of each token embedding.
103 pub dim: usize,
104 /// Number of bits per residual dimension (typically 2 or 4).
105 pub nbits: u32,
106 /// Number of coarse centroids (k in k-means).
107 pub k_centroids: usize,
108 /// Maximum iterations for k-means clustering.
109 pub max_kmeans_iters: usize,
110}
111
112/// A fully-built PLAID index over a corpus of multi-vector embeddings.
113///
114/// Stored in a **flat, StridedTensor-style layout** — every document's
115/// encoded tokens live in one of two big contiguous buffers,
116/// `doc_centroid_ids` (one u32 per token) and `doc_residual_bytes`
117/// (`packed_bytes_per_token` u8s per token). Per-document slicing
118/// happens through the precomputed `doc_offsets` cumulative lengths.
119///
120/// This layout matches fast-plaid's `StridedTensor` and buys three
121/// things over the older `Vec<Vec<EncodedVector>>` we used to keep:
122///
123/// 1. Search doesn't have to walk per-document `Vec`s on every query
124/// to gather codes and residuals — it slices into the flat buffer.
125/// 2. Peak RAM drops: 6.8M tokens at nbits=2 was ~500 MiB of heap
126/// for the old nested-Vec headers; the flat layout is ~240 MiB.
127/// 3. GPU decode can `Tensor::from_slice` the whole slice for a batch
128/// of candidates in one kernel launch instead of one-per-token.
129///
130/// Callers that want the per-token view use [`Index::doc_centroid_ids`]
131/// / [`Index::doc_residual_bytes`] — slices into the flat buffers with
132/// no allocation.
133#[derive(Debug, Clone)]
134pub struct Index {
135 pub params: IndexParams,
136 pub codec: ResidualCodec,
137 pub doc_ids: Vec<u64>,
138 /// Flat `[total_tokens]` vector of per-token centroid indices.
139 pub doc_centroid_ids: Vec<u32>,
140 /// Flat `[total_tokens * packed_bytes_per_token]` residual bytes,
141 /// row-major — each `packed_bytes_per_token`-long slice is one
142 /// token's packed residual.
143 pub doc_residual_bytes: Vec<u8>,
144 /// Cumulative per-document token counts; length `num_docs + 1`,
145 /// `doc_offsets[i + 1] - doc_offsets[i]` = `n_tokens` for doc `i`.
146 pub doc_offsets: Vec<usize>,
147 /// Centroid → tokens inverted file used for candidate generation.
148 pub ivf: InvertedFile,
149}
150
151impl Index {
152 /// Construct an `Index` from per-document `EncodedVector`s.
153 ///
154 /// Convenience for tests, [`crate::update::apply_update`], and
155 /// [`crate::persistence`]'s legacy-format loader: they still think
156 /// in terms of `Vec<Vec<EncodedVector>>`, and this flattens that
157 /// into the canonical [`Index`] layout in one pass.
158 ///
159 /// `ivf` should already reflect the `doc_tokens` contents — this
160 /// helper does not recompute the inverted file, only the flat
161 /// per-token buffers.
162 pub fn from_encoded_docs(
163 params: IndexParams,
164 codec: ResidualCodec,
165 doc_ids: Vec<u64>,
166 doc_tokens: Vec<Vec<EncodedVector>>,
167 ivf: InvertedFile,
168 ) -> Self {
169 let packed_bytes = codec.packed_bytes();
170 let total_tokens: usize = doc_tokens.iter().map(Vec::len).sum();
171 let mut doc_centroid_ids: Vec<u32> = Vec::with_capacity(total_tokens);
172 let mut doc_residual_bytes: Vec<u8> =
173 Vec::with_capacity(total_tokens * packed_bytes);
174 let mut doc_offsets: Vec<usize> =
175 Vec::with_capacity(doc_tokens.len() + 1);
176 doc_offsets.push(0);
177 for tokens in &doc_tokens {
178 for ev in tokens {
179 doc_centroid_ids.push(ev.centroid_id);
180 debug_assert_eq!(ev.codes.len(), packed_bytes);
181 doc_residual_bytes.extend_from_slice(&ev.codes);
182 }
183 doc_offsets.push(doc_centroid_ids.len());
184 }
185
186 Self {
187 params,
188 codec,
189 doc_ids,
190 doc_centroid_ids,
191 doc_residual_bytes,
192 doc_offsets,
193 ivf,
194 }
195 }
196
197 /// Number of documents currently stored in the index.
198 pub fn num_documents(&self) -> usize {
199 self.doc_ids.len()
200 }
201
202 /// Total number of encoded tokens across all documents.
203 pub fn num_tokens(&self) -> usize {
204 self.doc_centroid_ids.len()
205 }
206
207 /// Find the position of a document inside [`Index::doc_ids`].
208 pub fn position_of(&self, doc_id: u64) -> Option<usize> {
209 self.doc_ids.iter().position(|id| *id == doc_id)
210 }
211
212 /// Number of encoded tokens for the `idx`-th document.
213 pub fn doc_token_count(&self, idx: usize) -> usize {
214 self.doc_offsets[idx + 1] - self.doc_offsets[idx]
215 }
216
217 /// Slice of per-token centroid indices for the `idx`-th document.
218 pub fn doc_centroid_ids(&self, idx: usize) -> &[u32] {
219 &self.doc_centroid_ids[self.doc_offsets[idx]..self.doc_offsets[idx + 1]]
220 }
221
222 /// Slice of packed residual bytes for the `idx`-th document,
223 /// row-major `[n_tokens, packed_bytes_per_token]`.
224 pub fn doc_residual_bytes(&self, idx: usize) -> &[u8] {
225 let pb = self.codec.packed_bytes();
226 let start = self.doc_offsets[idx] * pb;
227 let end = self.doc_offsets[idx + 1] * pb;
228 &self.doc_residual_bytes[start..end]
229 }
230
231 /// Reconstruct the `idx`-th document's tokens as an owned
232 /// `Vec<EncodedVector>`.
233 ///
234 /// Allocation-heavy — prefer [`Index::doc_centroid_ids`] /
235 /// [`Index::doc_residual_bytes`] on the hot path. Provided as a
236 /// convenience for [`crate::update::apply_update`] and other
237 /// callers that work in `EncodedVector` terms.
238 pub fn doc_tokens_vec(&self, idx: usize) -> Vec<EncodedVector> {
239 let pb = self.codec.packed_bytes();
240 let cids = self.doc_centroid_ids(idx);
241 let res = self.doc_residual_bytes(idx);
242 (0..cids.len())
243 .map(|i| EncodedVector {
244 centroid_id: cids[i],
245 codes: res[i * pb..(i + 1) * pb].to_vec(),
246 })
247 .collect()
248 }
249}
250
251/// Build a [`Index`] from a corpus of documents.
252///
253/// Every document must share the same embedding dimensionality as
254/// `params.dim`. Documents with zero tokens are preserved in the index —
255/// they contribute nothing to centroid/codec training but still occupy a
256/// slot in `doc_ids` so callers can resolve their position by `doc_id`
257/// later.
258///
259/// # Errors
260///
261/// Returns [`PlaidError::Tensor`] if the matmul-driven k-means
262/// training or nearest-centroid assignment fails.
263///
264/// # Panics
265///
266/// Panics if any document's flat length is not a multiple of `dim`, if
267/// the total number of tokens is smaller than `params.k_centroids`, or
268/// if `params.k_centroids == 0`.
269///
270/// [`PlaidError::Tensor`]: crate::PlaidError::Tensor
271pub fn build_index(
272 documents: &[DocumentTokens],
273 params: IndexParams,
274) -> Result<Index> {
275 assert!(params.dim > 0, "build_index: dim must be positive");
276 assert!(
277 params.k_centroids > 0,
278 "build_index: k_centroids must be positive"
279 );
280 assert!(
281 params.nbits > 0 && params.nbits <= 8,
282 "build_index: nbits must be in 1..=8, got {}",
283 params.nbits,
284 );
285
286 for doc in documents {
287 assert!(
288 doc.tokens.len() == doc.n_tokens * params.dim,
289 "build_index: doc {} declared {} tokens but carries {} f32s (dim={})",
290 doc.doc_id,
291 doc.n_tokens,
292 doc.tokens.len(),
293 params.dim,
294 );
295 }
296
297 // Flatten all token embeddings into one training cloud and forward
298 // to the pool-based core path. This mirrors the memory-lean path
299 // used by `build_index_from_pool` callers that already own a
300 // contiguous buffer.
301 let total_tokens: usize = documents.iter().map(|d| d.n_tokens).sum();
302 let mut pool: Vec<f32> = Vec::with_capacity(total_tokens * params.dim);
303 let mut doc_meta: Vec<(u64, usize)> = Vec::with_capacity(documents.len());
304 for doc in documents {
305 pool.extend_from_slice(&doc.tokens);
306 doc_meta.push((doc.doc_id, doc.n_tokens));
307 }
308
309 build_index_from_pool(pool, doc_meta, params)
310}
311
312/// Build an [`Index`] from a pre-assembled token pool and per-document
313/// `(doc_id, n_tokens)` metadata.
314///
315/// This is the memory-lean entry point: callers that can stream tokens
316/// straight into a single contiguous `Vec<f32>` (e.g. the `EmbeddingDb`
317/// bridge) avoid holding both a `Vec<DocumentTokens>` *and* a flat pool
318/// at the same time — on a real corpus that doubling is worth several
319/// GB of peak RSS.
320///
321/// `pool` is laid out row-major with `n_tokens × dim` entries, where
322/// `n_tokens = doc_meta.iter().map(|(_, n)| n).sum()`. Documents keep
323/// the order of `doc_meta` in the resulting index.
324///
325/// # Errors
326///
327/// Returns [`PlaidError::Tensor`] if the matmul-driven k-means
328/// training or nearest-centroid assignment fails.
329///
330/// # Panics
331///
332/// Panics if `pool.len()` is not a multiple of `dim`, if the sum of
333/// `doc_meta`'s token counts disagrees with `pool.len() / dim`, if the
334/// total number of tokens is smaller than `params.k_centroids`, or if
335/// `params.k_centroids == 0`.
336///
337/// [`PlaidError::Tensor`]: crate::PlaidError::Tensor
338pub fn build_index_from_pool(
339 pool: Vec<f32>,
340 doc_meta: Vec<(u64, usize)>,
341 params: IndexParams,
342) -> Result<Index> {
343 assert!(
344 params.dim > 0,
345 "build_index_from_pool: dim must be positive"
346 );
347 assert!(
348 params.k_centroids > 0,
349 "build_index_from_pool: k_centroids must be positive"
350 );
351 assert!(
352 params.nbits > 0 && params.nbits <= 8,
353 "build_index_from_pool: nbits must be in 1..=8, got {}",
354 params.nbits,
355 );
356 assert!(
357 pool.len().is_multiple_of(params.dim),
358 "build_index_from_pool: pool length {} is not a multiple of dim {}",
359 pool.len(),
360 params.dim,
361 );
362 let total_tokens = pool.len() / params.dim;
363 let meta_tokens: usize = doc_meta.iter().map(|(_, n)| n).sum();
364 assert_eq!(
365 total_tokens, meta_tokens,
366 "build_index_from_pool: pool carries {total_tokens} tokens but doc_meta sums to {meta_tokens}",
367 );
368 assert!(
369 total_tokens >= params.k_centroids,
370 "build_index_from_pool: need at least {} tokens for {} centroids, got {}",
371 params.k_centroids,
372 params.k_centroids,
373 total_tokens,
374 );
375
376 // Upload the [n, dim] pool tensor once and share it across the
377 // k-means phase and the final batch-encode pass. Without this, each
378 // phase uploaded its own ~3.47 GB copy for a real docbert corpus,
379 // cudarc's caching allocator held the stale block, and the PLAID
380 // build tipped a 12 GB card into CUDA OOM as soon as the encoder
381 // model was also resident.
382 let pool_bytes = pool.len() * std::mem::size_of::<f32>();
383 let device_info = crate::device::device_memory_info();
384 eprintln!(
385 " Pool: {total_tokens} tokens × {} dim = {} MiB{}",
386 params.dim,
387 pool_bytes / (1 << 20),
388 match device_info {
389 Some((free, total)) => format!(
390 " (device: {} MiB free / {} MiB total)",
391 free / (1 << 20),
392 total / (1 << 20),
393 ),
394 None => String::new(),
395 },
396 );
397
398 // 1. Train coarse centroids with k-means. `fit` subsamples to
399 // `k * MAX_POINTS_PER_CENTROID` rows and uploads only that
400 // subsample, so peak VRAM here is bounded regardless of pool
401 // size or `dim`.
402 let centroids = fit(
403 &pool,
404 params.k_centroids,
405 params.dim,
406 params.max_kmeans_iters.max(1),
407 )?;
408
409 // 2. Residuals for quantizer training. We only need enough samples
410 // to place `2^nbits` quantile cutoffs — materialising a
411 // residual-per-token for a large corpus would allocate multiple
412 // gigabytes for no statistical benefit. Stride-sample the pool,
413 // compute residuals only for the sampled tokens, and move on.
414 let sample_stride =
415 total_tokens.div_ceil(MAX_QUANTIZER_TRAINING_TOKENS).max(1);
416 let sample_count = total_tokens.div_ceil(sample_stride);
417
418 let mut sampled_tokens: Vec<f32> =
419 Vec::with_capacity(sample_count * params.dim);
420 for (i, token) in pool.chunks_exact(params.dim).enumerate() {
421 if i.is_multiple_of(sample_stride) {
422 sampled_tokens.extend_from_slice(token);
423 }
424 }
425 let sample_assignments =
426 assign_points(&sampled_tokens, ¢roids, params.dim)?;
427 let mut residual_sample: Vec<f32> =
428 Vec::with_capacity(sampled_tokens.len());
429 for (token, &cluster) in sampled_tokens
430 .chunks_exact(params.dim)
431 .zip(&sample_assignments)
432 {
433 let centroid =
434 ¢roids[cluster * params.dim..(cluster + 1) * params.dim];
435 for (t, c) in token.iter().zip(centroid) {
436 residual_sample.push(*t - *c);
437 }
438 }
439 drop(sampled_tokens);
440
441 // 3. Learn cutoffs + weights on the sampled residuals.
442 let (bucket_cutoffs, bucket_weights) =
443 train_quantizer(residual_sample, params.nbits);
444
445 let codec = ResidualCodec {
446 nbits: params.nbits,
447 dim: params.dim,
448 centroids,
449 bucket_cutoffs,
450 bucket_weights,
451 };
452 codec.validate()?;
453
454 // 4. Encode every token across the whole corpus. The encoder
455 // walks the host `pool` in `[chunk_rows, dim]` tiles, uploads
456 // each tile, runs assign + bucketize + pack on it, drains the
457 // results back to host, and drops the tile before the next
458 // one uploads. Peak VRAM is bounded by
459 // `codec state + one chunk` (≈128 MiB at default settings) —
460 // unchanged by corpus size or embedding dimension, so a
461 // 1536-dim model on a 6.8M-token corpus builds without the
462 // 40 GB contiguous allocation the old single-shot path
463 // required.
464 let (doc_centroid_ids, doc_residual_bytes) =
465 codec.batch_encode_tokens(&pool)?;
466 drop(pool);
467 debug_assert_eq!(doc_centroid_ids.len(), total_tokens);
468 debug_assert_eq!(
469 doc_residual_bytes.len(),
470 total_tokens * codec.packed_bytes()
471 );
472
473 let mut doc_ids = Vec::with_capacity(doc_meta.len());
474 let mut doc_offsets = Vec::with_capacity(doc_meta.len() + 1);
475 doc_offsets.push(0usize);
476 let mut running = 0usize;
477 for (doc_id, n_tok) in doc_meta {
478 doc_ids.push(doc_id);
479 running += n_tok;
480 doc_offsets.push(running);
481 }
482
483 let ivf = build_inverted_file_from_flat(
484 &doc_centroid_ids,
485 &doc_offsets,
486 params.k_centroids,
487 );
488
489 Ok(Index {
490 params,
491 codec,
492 doc_ids,
493 doc_centroid_ids,
494 doc_residual_bytes,
495 doc_offsets,
496 ivf,
497 })
498}
499
500/// Build the centroid → unique-doc-ids inverted file from a flat-layout
501/// index's centroid assignments.
502///
503/// Each doc contributes at most one entry per centroid it touches;
504/// entries within a list are sorted ascending. `doc_offsets` carries
505/// the cumulative token counts (length `num_docs + 1`) that delimit
506/// each document's slice of `doc_centroid_ids`.
507pub(crate) fn build_inverted_file_from_flat(
508 doc_centroid_ids: &[u32],
509 doc_offsets: &[usize],
510 k_centroids: usize,
511) -> InvertedFile {
512 let mut lists: Vec<Vec<u32>> = vec![Vec::new(); k_centroids];
513 let n_docs = doc_offsets.len().saturating_sub(1);
514 // Per-centroid stamp of the last doc that touched it, so each doc
515 // contributes one entry per distinct centroid without the per-doc
516 // sort + dedup + allocation of the naive version — that tripled
517 // index load time on a 20M-token corpus. Pushing in doc order
518 // keeps every list sorted ascending, same as before.
519 let mut last_doc: Vec<u32> = vec![u32::MAX; k_centroids];
520 for doc_idx in 0..n_docs {
521 let doc_codes =
522 &doc_centroid_ids[doc_offsets[doc_idx]..doc_offsets[doc_idx + 1]];
523 for &cid in doc_codes {
524 if last_doc[cid as usize] != doc_idx as u32 {
525 last_doc[cid as usize] = doc_idx as u32;
526 lists[cid as usize].push(doc_idx as u32);
527 }
528 }
529 }
530 InvertedFile { lists }
531}
532
533#[cfg(test)]
534mod tests {
535 use super::*;
536 use crate::distance::squared_l2;
537
538 /// Build a tiny 2-D corpus with two clear clusters of tokens.
539 fn small_corpus() -> Vec<DocumentTokens> {
540 vec![
541 DocumentTokens {
542 doc_id: 1,
543 tokens: vec![0.0, 0.0, 0.1, 0.2, -0.1, 0.1],
544 n_tokens: 3,
545 },
546 DocumentTokens {
547 doc_id: 2,
548 tokens: vec![10.0, 10.0, 10.2, 9.9, 9.8, 10.1],
549 n_tokens: 3,
550 },
551 DocumentTokens {
552 doc_id: 3,
553 tokens: vec![0.3, -0.2, 9.7, 10.2],
554 n_tokens: 2,
555 },
556 ]
557 }
558
559 fn default_params() -> IndexParams {
560 IndexParams {
561 dim: 2,
562 nbits: 2,
563 k_centroids: 2,
564 max_kmeans_iters: 50,
565 }
566 }
567
568 #[test]
569 fn build_index_encodes_every_token() {
570 let docs = small_corpus();
571 let params = default_params();
572 let expected_total: usize = docs.iter().map(|d| d.n_tokens).sum();
573
574 let index = build_index(&docs, params).unwrap();
575
576 assert_eq!(index.num_documents(), docs.len());
577 assert_eq!(index.num_tokens(), expected_total);
578 for (i, doc) in docs.iter().enumerate() {
579 assert_eq!(index.doc_token_count(i), doc.n_tokens);
580 }
581 }
582
583 #[test]
584 fn build_index_assigns_tokens_in_each_cluster_to_the_closest_centroid() {
585 let docs = small_corpus();
586 let params = default_params();
587 let index = build_index(&docs, params).unwrap();
588
589 // The two tight clusters around (0,0) and (10,10) should produce
590 // centroids close to those means.
591 let c0 = &index.codec.centroids[0..2];
592 let c1 = &index.codec.centroids[2..4];
593
594 let (near_origin, near_ten) =
595 if squared_l2(c0, &[0.0, 0.0]) < squared_l2(c0, &[10.0, 10.0]) {
596 (c0, c1)
597 } else {
598 (c1, c0)
599 };
600
601 assert!(squared_l2(near_origin, &[0.0, 0.0]) < 1.0);
602 assert!(squared_l2(near_ten, &[10.0, 10.0]) < 1.0);
603 }
604
605 #[test]
606 fn build_index_round_trip_reconstruction_error_is_bounded() {
607 let docs = small_corpus();
608 let params = default_params();
609 let index = build_index(&docs, params).unwrap();
610
611 for (i, doc) in docs.iter().enumerate() {
612 let encoded_doc = index.doc_tokens_vec(i);
613 for (token, encoded) in
614 doc.tokens.chunks_exact(params.dim).zip(encoded_doc.iter())
615 {
616 let decoded = index.codec.decode_vector(encoded).unwrap();
617 let err = squared_l2(token, &decoded).sqrt();
618 // Residuals on a 2-D toy corpus with tight clusters stay
619 // small; each bucket should cover well under 0.5 per dim.
620 assert!(
621 err < 0.6,
622 "reconstruction error {err} too large for token {token:?}"
623 );
624 }
625 }
626 }
627
628 #[test]
629 fn build_index_preserves_document_id_order() {
630 let docs = small_corpus();
631 let index = build_index(&docs, default_params()).unwrap();
632 let expected_ids: Vec<u64> = docs.iter().map(|d| d.doc_id).collect();
633 assert_eq!(index.doc_ids, expected_ids);
634 assert_eq!(index.position_of(2), Some(1));
635 assert_eq!(index.position_of(999), None);
636 }
637
638 #[test]
639 fn build_index_handles_document_with_no_tokens() {
640 // Empty documents are still indexable: they contribute no tokens
641 // to training but keep their slot so callers can look them up.
642 let mut docs = small_corpus();
643 docs.push(DocumentTokens {
644 doc_id: 42,
645 tokens: vec![],
646 n_tokens: 0,
647 });
648 let index = build_index(&docs, default_params()).unwrap();
649 assert_eq!(index.num_documents(), 4);
650 assert_eq!(index.doc_token_count(3), 0);
651 }
652
653 #[test]
654 #[should_panic(expected = "declared")]
655 fn build_index_panics_on_mismatched_token_count() {
656 let docs = vec![DocumentTokens {
657 doc_id: 1,
658 tokens: vec![0.0, 0.0, 1.0],
659 n_tokens: 2, // says 2 but only 3 f32s and dim=2
660 }];
661 let _ = build_index(&docs, default_params()).unwrap();
662 }
663
664 #[test]
665 fn build_index_ivf_has_one_list_per_centroid() {
666 let docs = small_corpus();
667 let index = build_index(&docs, default_params()).unwrap();
668
669 assert_eq!(
670 index.ivf.num_centroids(),
671 default_params().k_centroids,
672 "IVF has one list per centroid",
673 );
674 }
675
676 #[test]
677 fn build_index_ivf_postings_cover_every_doc_that_touches_each_centroid() {
678 // Every (doc, token) pair implies the doc_idx must appear in
679 // that centroid's posting list. This is the PLAID "centroid →
680 // unique passage ids" contract: postings are indexed by doc.
681 let docs = small_corpus();
682 let index = build_index(&docs, default_params()).unwrap();
683
684 for doc_idx in 0..index.num_documents() {
685 for &cid in index.doc_centroid_ids(doc_idx) {
686 let postings = index.ivf.docs_for_centroid(cid as usize);
687 assert!(
688 postings.contains(&(doc_idx as u32)),
689 "doc_idx={doc_idx} missing from centroid {cid} postings",
690 );
691 }
692 }
693 }
694
695 #[test]
696 fn build_index_ivf_postings_are_unique_per_centroid() {
697 // PLAID stores centroid → unique doc ids, not token refs. A doc
698 // with multiple tokens in the same centroid must appear at most
699 // once in that centroid's list.
700 let docs = small_corpus();
701 let index = build_index(&docs, default_params()).unwrap();
702
703 for c in 0..index.ivf.num_centroids() {
704 let postings = index.ivf.docs_for_centroid(c);
705 let mut unique: Vec<u32> = postings.to_vec();
706 unique.sort_unstable();
707 unique.dedup();
708 assert_eq!(
709 unique.len(),
710 postings.len(),
711 "centroid {c} has duplicate doc entries: {postings:?}",
712 );
713 }
714 }
715
716 #[test]
717 fn build_index_dedupes_repeated_tokens_in_same_centroid() {
718 // Doc 1 has three tokens that all cluster to the same coarse
719 // centroid. The doc_idx should show up once in that centroid's
720 // posting, not three times.
721 let docs = vec![
722 DocumentTokens {
723 doc_id: 1,
724 tokens: vec![0.0, 0.0, 0.05, -0.02, -0.03, 0.01],
725 n_tokens: 3,
726 },
727 DocumentTokens {
728 doc_id: 2,
729 tokens: vec![10.0, 10.0, 10.1, 9.9],
730 n_tokens: 2,
731 },
732 ];
733 let index = build_index(&docs, default_params()).unwrap();
734
735 for c in 0..index.ivf.num_centroids() {
736 let postings = index.ivf.docs_for_centroid(c);
737 let count_of_doc_0 = postings.iter().filter(|&&d| d == 0).count();
738 assert!(
739 count_of_doc_0 <= 1,
740 "doc 0 appears {count_of_doc_0} times in centroid {c}",
741 );
742 }
743 }
744
745 #[test]
746 fn inverted_file_out_of_range_returns_empty_slice() {
747 let ivf = InvertedFile {
748 lists: vec![vec![0u32]],
749 };
750 assert_eq!(ivf.docs_for_centroid(0).len(), 1);
751 assert!(ivf.docs_for_centroid(999).is_empty());
752 }
753
754 #[test]
755 #[should_panic(expected = "at least")]
756 fn build_index_panics_when_too_few_tokens_for_k() {
757 let docs = vec![DocumentTokens {
758 doc_id: 1,
759 tokens: vec![0.0, 1.0],
760 n_tokens: 1,
761 }];
762 let params = IndexParams {
763 dim: 2,
764 nbits: 2,
765 k_centroids: 4,
766 max_kmeans_iters: 10,
767 };
768 let _ = build_index(&docs, params).unwrap();
769 }
770}