Skip to main content

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, &centroids, 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            &centroids[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}