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 /// compatibility shim for [`crate::update::apply_update`] and
237 /// legacy callers that still 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 for doc_idx in 0..n_docs {
515 let doc_codes =
516 &doc_centroid_ids[doc_offsets[doc_idx]..doc_offsets[doc_idx + 1]];
517 let mut touched: Vec<u32> = doc_codes.to_vec();
518 touched.sort_unstable();
519 touched.dedup();
520 for cid in touched {
521 lists[cid as usize].push(doc_idx as u32);
522 }
523 }
524 InvertedFile { lists }
525}
526
527#[cfg(test)]
528mod tests {
529 use super::*;
530 use crate::distance::squared_l2;
531
532 /// Build a tiny 2-D corpus with two clear clusters of tokens.
533 fn small_corpus() -> Vec<DocumentTokens> {
534 vec![
535 DocumentTokens {
536 doc_id: 1,
537 tokens: vec![0.0, 0.0, 0.1, 0.2, -0.1, 0.1],
538 n_tokens: 3,
539 },
540 DocumentTokens {
541 doc_id: 2,
542 tokens: vec![10.0, 10.0, 10.2, 9.9, 9.8, 10.1],
543 n_tokens: 3,
544 },
545 DocumentTokens {
546 doc_id: 3,
547 tokens: vec![0.3, -0.2, 9.7, 10.2],
548 n_tokens: 2,
549 },
550 ]
551 }
552
553 fn default_params() -> IndexParams {
554 IndexParams {
555 dim: 2,
556 nbits: 2,
557 k_centroids: 2,
558 max_kmeans_iters: 50,
559 }
560 }
561
562 #[test]
563 fn build_index_encodes_every_token() {
564 let docs = small_corpus();
565 let params = default_params();
566 let expected_total: usize = docs.iter().map(|d| d.n_tokens).sum();
567
568 let index = build_index(&docs, params).unwrap();
569
570 assert_eq!(index.num_documents(), docs.len());
571 assert_eq!(index.num_tokens(), expected_total);
572 for (i, doc) in docs.iter().enumerate() {
573 assert_eq!(index.doc_token_count(i), doc.n_tokens);
574 }
575 }
576
577 #[test]
578 fn build_index_assigns_tokens_in_each_cluster_to_the_closest_centroid() {
579 let docs = small_corpus();
580 let params = default_params();
581 let index = build_index(&docs, params).unwrap();
582
583 // The two tight clusters around (0,0) and (10,10) should produce
584 // centroids close to those means.
585 let c0 = &index.codec.centroids[0..2];
586 let c1 = &index.codec.centroids[2..4];
587
588 let (near_origin, near_ten) =
589 if squared_l2(c0, &[0.0, 0.0]) < squared_l2(c0, &[10.0, 10.0]) {
590 (c0, c1)
591 } else {
592 (c1, c0)
593 };
594
595 assert!(squared_l2(near_origin, &[0.0, 0.0]) < 1.0);
596 assert!(squared_l2(near_ten, &[10.0, 10.0]) < 1.0);
597 }
598
599 #[test]
600 fn build_index_round_trip_reconstruction_error_is_bounded() {
601 let docs = small_corpus();
602 let params = default_params();
603 let index = build_index(&docs, params).unwrap();
604
605 for (i, doc) in docs.iter().enumerate() {
606 let encoded_doc = index.doc_tokens_vec(i);
607 for (token, encoded) in
608 doc.tokens.chunks_exact(params.dim).zip(encoded_doc.iter())
609 {
610 let decoded = index.codec.decode_vector(encoded).unwrap();
611 let err = squared_l2(token, &decoded).sqrt();
612 // Residuals on a 2-D toy corpus with tight clusters stay
613 // small; each bucket should cover well under 0.5 per dim.
614 assert!(
615 err < 0.6,
616 "reconstruction error {err} too large for token {token:?}"
617 );
618 }
619 }
620 }
621
622 #[test]
623 fn build_index_preserves_document_id_order() {
624 let docs = small_corpus();
625 let index = build_index(&docs, default_params()).unwrap();
626 let expected_ids: Vec<u64> = docs.iter().map(|d| d.doc_id).collect();
627 assert_eq!(index.doc_ids, expected_ids);
628 assert_eq!(index.position_of(2), Some(1));
629 assert_eq!(index.position_of(999), None);
630 }
631
632 #[test]
633 fn build_index_handles_document_with_no_tokens() {
634 // Empty documents are still indexable: they contribute no tokens
635 // to training but keep their slot so callers can look them up.
636 let mut docs = small_corpus();
637 docs.push(DocumentTokens {
638 doc_id: 42,
639 tokens: vec![],
640 n_tokens: 0,
641 });
642 let index = build_index(&docs, default_params()).unwrap();
643 assert_eq!(index.num_documents(), 4);
644 assert_eq!(index.doc_token_count(3), 0);
645 }
646
647 #[test]
648 #[should_panic(expected = "declared")]
649 fn build_index_panics_on_mismatched_token_count() {
650 let docs = vec![DocumentTokens {
651 doc_id: 1,
652 tokens: vec![0.0, 0.0, 1.0],
653 n_tokens: 2, // says 2 but only 3 f32s and dim=2
654 }];
655 let _ = build_index(&docs, default_params()).unwrap();
656 }
657
658 #[test]
659 fn build_index_ivf_has_one_list_per_centroid() {
660 let docs = small_corpus();
661 let index = build_index(&docs, default_params()).unwrap();
662
663 assert_eq!(
664 index.ivf.num_centroids(),
665 default_params().k_centroids,
666 "IVF has one list per centroid",
667 );
668 }
669
670 #[test]
671 fn build_index_ivf_postings_cover_every_doc_that_touches_each_centroid() {
672 // Every (doc, token) pair implies the doc_idx must appear in
673 // that centroid's posting list. This is the PLAID "centroid →
674 // unique passage ids" contract: postings are indexed by doc.
675 let docs = small_corpus();
676 let index = build_index(&docs, default_params()).unwrap();
677
678 for doc_idx in 0..index.num_documents() {
679 for &cid in index.doc_centroid_ids(doc_idx) {
680 let postings = index.ivf.docs_for_centroid(cid as usize);
681 assert!(
682 postings.contains(&(doc_idx as u32)),
683 "doc_idx={doc_idx} missing from centroid {cid} postings",
684 );
685 }
686 }
687 }
688
689 #[test]
690 fn build_index_ivf_postings_are_unique_per_centroid() {
691 // PLAID stores centroid → unique doc ids, not token refs. A doc
692 // with multiple tokens in the same centroid must appear at most
693 // once in that centroid's list.
694 let docs = small_corpus();
695 let index = build_index(&docs, default_params()).unwrap();
696
697 for c in 0..index.ivf.num_centroids() {
698 let postings = index.ivf.docs_for_centroid(c);
699 let mut unique: Vec<u32> = postings.to_vec();
700 unique.sort_unstable();
701 unique.dedup();
702 assert_eq!(
703 unique.len(),
704 postings.len(),
705 "centroid {c} has duplicate doc entries: {postings:?}",
706 );
707 }
708 }
709
710 #[test]
711 fn build_index_dedupes_repeated_tokens_in_same_centroid() {
712 // Doc 1 has three tokens that all cluster to the same coarse
713 // centroid. The doc_idx should show up once in that centroid's
714 // posting, not three times.
715 let docs = vec![
716 DocumentTokens {
717 doc_id: 1,
718 tokens: vec![0.0, 0.0, 0.05, -0.02, -0.03, 0.01],
719 n_tokens: 3,
720 },
721 DocumentTokens {
722 doc_id: 2,
723 tokens: vec![10.0, 10.0, 10.1, 9.9],
724 n_tokens: 2,
725 },
726 ];
727 let index = build_index(&docs, default_params()).unwrap();
728
729 for c in 0..index.ivf.num_centroids() {
730 let postings = index.ivf.docs_for_centroid(c);
731 let count_of_doc_0 = postings.iter().filter(|&&d| d == 0).count();
732 assert!(
733 count_of_doc_0 <= 1,
734 "doc 0 appears {count_of_doc_0} times in centroid {c}",
735 );
736 }
737 }
738
739 #[test]
740 fn inverted_file_out_of_range_returns_empty_slice() {
741 let ivf = InvertedFile {
742 lists: vec![vec![0u32]],
743 };
744 assert_eq!(ivf.docs_for_centroid(0).len(), 1);
745 assert!(ivf.docs_for_centroid(999).is_empty());
746 }
747
748 #[test]
749 #[should_panic(expected = "at least")]
750 fn build_index_panics_when_too_few_tokens_for_k() {
751 let docs = vec![DocumentTokens {
752 doc_id: 1,
753 tokens: vec![0.0, 1.0],
754 n_tokens: 1,
755 }];
756 let params = IndexParams {
757 dim: 2,
758 nbits: 2,
759 k_centroids: 4,
760 max_kmeans_iters: 10,
761 };
762 let _ = build_index(&docs, params).unwrap();
763 }
764}