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    /// 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, &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    // 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}