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/// Inverted file: for each centroid, the sorted list of unique document
24/// indices that have at least one token clustered in that centroid.
25///
26/// This mirrors PLAID's "centroid → unique passage ids" layout from
27/// §3 of the paper: candidate generation only needs to know which
28/// documents are reachable via a probed centroid, and deduplicating
29/// per-doc keeps the posting lists small even when a single document
30/// has many tokens mapped to the same cluster.
31///
32/// The search path uses this to expand a query token to a shortlist of
33/// document candidates: find the centroids with the highest dot-product
34/// against the query token, then gather every document listed under
35/// those centroids.
36#[derive(Debug, Clone, Default)]
37pub struct InvertedFile {
38    /// `lists[c]` holds the sorted, deduplicated `doc_idx`s of every
39    /// document with at least one token assigned to centroid `c`.
40    pub lists: Vec<Vec<u32>>,
41}
42
43impl InvertedFile {
44    /// Total number of centroids the IVF spans.
45    pub fn num_centroids(&self) -> usize {
46        self.lists.len()
47    }
48
49    /// Document indices currently associated with `centroid_id`, or an
50    /// empty slice if the centroid is out of range. Entries are sorted
51    /// ascending and contain no duplicates.
52    pub fn docs_for_centroid(&self, centroid_id: usize) -> &[u32] {
53        self.lists
54            .get(centroid_id)
55            .map(Vec::as_slice)
56            .unwrap_or(&[])
57    }
58
59    /// Total number of (centroid, doc) postings across every list.
60    ///
61    /// This is the sum of `lists[c].len()` over all centroids `c`. It
62    /// is at most `num_centroids * num_documents` and at least equal
63    /// to the number of documents that contain any tokens at all.
64    pub fn total_doc_postings(&self) -> usize {
65        self.lists.iter().map(Vec::len).sum()
66    }
67}
68
69/// A single document's worth of token embeddings, ready to index.
70///
71/// `tokens` is a flat row-major `n_tokens × dim` buffer. Keeping the
72/// tokens flat mirrors the way docbert already stores ColBERT outputs in
73/// `embeddings.db` and avoids an intermediate `Vec<Vec<f32>>` allocation.
74#[derive(Debug, Clone)]
75pub struct DocumentTokens {
76    pub doc_id: u64,
77    pub tokens: Vec<f32>,
78    pub n_tokens: usize,
79}
80
81impl DocumentTokens {
82    /// Total number of f32 values this document contributes.
83    pub fn flat_len(&self) -> usize {
84        self.tokens.len()
85    }
86}
87
88/// Parameters that control how an [`Index`] is built.
89#[derive(Debug, Clone, Copy)]
90pub struct IndexParams {
91    /// Dimensionality of each token embedding.
92    pub dim: usize,
93    /// Number of bits per residual dimension (typically 2 or 4).
94    pub nbits: u32,
95    /// Number of coarse centroids (k in k-means).
96    pub k_centroids: usize,
97    /// Maximum iterations for k-means clustering.
98    pub max_kmeans_iters: usize,
99}
100
101/// A fully-built PLAID index over a corpus of multi-vector embeddings.
102///
103/// Holds the trained codec, the encoded token embeddings of every
104/// document, and an inverted file mapping centroids back to the tokens
105/// clustered in them. Future layers (search) will read from this state
106/// directly without mutating it.
107#[derive(Debug, Clone)]
108pub struct Index {
109    pub params: IndexParams,
110    pub codec: ResidualCodec,
111    pub doc_ids: Vec<u64>,
112    /// `doc_tokens[i]` is the encoded token sequence of the i-th document.
113    pub doc_tokens: Vec<Vec<EncodedVector>>,
114    /// Centroid → tokens inverted file used for candidate generation.
115    pub ivf: InvertedFile,
116}
117
118impl Index {
119    /// Number of documents currently stored in the index.
120    pub fn num_documents(&self) -> usize {
121        self.doc_ids.len()
122    }
123
124    /// Total number of encoded tokens across all documents.
125    pub fn num_tokens(&self) -> usize {
126        self.doc_tokens.iter().map(Vec::len).sum()
127    }
128
129    /// Find the position of a document inside [`Index::doc_ids`].
130    pub fn position_of(&self, doc_id: u64) -> Option<usize> {
131        self.doc_ids.iter().position(|id| *id == doc_id)
132    }
133}
134
135/// Build a [`Index`] from a corpus of documents.
136///
137/// Every document must share the same embedding dimensionality as
138/// `params.dim`. Documents with zero tokens are preserved in the index —
139/// they contribute nothing to centroid/codec training but still occupy a
140/// slot in `doc_ids` so callers can resolve their position by `doc_id`
141/// later.
142///
143/// # Errors
144///
145/// Returns [`PlaidError::Tensor`] if the matmul-driven k-means
146/// training or nearest-centroid assignment fails.
147///
148/// # Panics
149///
150/// Panics if any document's flat length is not a multiple of `dim`, if
151/// the total number of tokens is smaller than `params.k_centroids`, or
152/// if `params.k_centroids == 0`.
153///
154/// [`PlaidError::Tensor`]: crate::PlaidError::Tensor
155pub fn build_index(
156    documents: &[DocumentTokens],
157    params: IndexParams,
158) -> Result<Index> {
159    assert!(params.dim > 0, "build_index: dim must be positive");
160    assert!(
161        params.k_centroids > 0,
162        "build_index: k_centroids must be positive"
163    );
164    assert!(
165        params.nbits > 0 && params.nbits <= 8,
166        "build_index: nbits must be in 1..=8, got {}",
167        params.nbits,
168    );
169
170    for doc in documents {
171        assert!(
172            doc.tokens.len() == doc.n_tokens * params.dim,
173            "build_index: doc {} declared {} tokens but carries {} f32s (dim={})",
174            doc.doc_id,
175            doc.n_tokens,
176            doc.tokens.len(),
177            params.dim,
178        );
179    }
180
181    // Flatten all token embeddings into one training cloud.
182    let total_tokens: usize = documents.iter().map(|d| d.n_tokens).sum();
183    assert!(
184        total_tokens >= params.k_centroids,
185        "build_index: need at least {} tokens for {} centroids, got {}",
186        params.k_centroids,
187        params.k_centroids,
188        total_tokens,
189    );
190    let mut pool: Vec<f32> = Vec::with_capacity(total_tokens * params.dim);
191    for doc in documents {
192        pool.extend_from_slice(&doc.tokens);
193    }
194
195    // 1. Train coarse centroids with k-means.
196    let centroids = fit(
197        &pool,
198        params.k_centroids,
199        params.dim,
200        params.max_kmeans_iters.max(1),
201    )?;
202
203    // 2. Compute residuals against trained centroids so we can train the
204    //    quantizer on a representative sample of errors.
205    let assignments = assign_points(&pool, &centroids, params.dim)?;
206    let mut residual_sample: Vec<f32> = Vec::with_capacity(pool.len());
207    for (token, &cluster) in pool.chunks_exact(params.dim).zip(&assignments) {
208        let centroid =
209            &centroids[cluster * params.dim..(cluster + 1) * params.dim];
210        for (t, c) in token.iter().zip(centroid) {
211            residual_sample.push(*t - *c);
212        }
213    }
214
215    // 3. Learn cutoffs + weights on observed residuals.
216    let (bucket_cutoffs, bucket_weights) =
217        train_quantizer(&residual_sample, params.nbits);
218
219    let codec = ResidualCodec {
220        nbits: params.nbits,
221        dim: params.dim,
222        centroids,
223        bucket_cutoffs,
224        bucket_weights,
225    };
226    codec.validate()?;
227
228    // 4. Encode every token across the whole corpus in one batched pass
229    //    (single matmul-driven nearest-centroid lookup), then split the
230    //    flat result back into per-document EncodedVectors and populate
231    //    the centroid → tokens inverted file along the way.
232    let (all_centroid_ids, all_codes) = codec.batch_encode_tokens(&pool)?;
233    let packed_per_token = codec.packed_bytes();
234
235    let mut doc_ids = Vec::with_capacity(documents.len());
236    let mut doc_tokens = Vec::with_capacity(documents.len());
237    let mut token_offset = 0usize;
238    for doc in documents {
239        doc_ids.push(doc.doc_id);
240        let n_tok = doc.n_tokens;
241        let cids = &all_centroid_ids[token_offset..token_offset + n_tok];
242        let codes_slice = &all_codes[token_offset * packed_per_token
243            ..(token_offset + n_tok) * packed_per_token];
244        let encoded: Vec<EncodedVector> = (0..n_tok)
245            .map(|i| EncodedVector {
246                centroid_id: cids[i],
247                codes: codes_slice
248                    [i * packed_per_token..(i + 1) * packed_per_token]
249                    .to_vec(),
250            })
251            .collect();
252        doc_tokens.push(encoded);
253        token_offset += n_tok;
254    }
255
256    let ivf = build_inverted_file(&doc_tokens, params.k_centroids);
257
258    Ok(Index {
259        params,
260        codec,
261        doc_ids,
262        doc_tokens,
263        ivf,
264    })
265}
266
267/// Build the centroid → unique-doc-ids inverted file from the encoded
268/// corpus. Each doc contributes at most one entry per centroid it
269/// touches; entries within a list are sorted ascending.
270pub(crate) fn build_inverted_file(
271    doc_tokens: &[Vec<EncodedVector>],
272    k_centroids: usize,
273) -> InvertedFile {
274    let mut lists: Vec<Vec<u32>> = vec![Vec::new(); k_centroids];
275    for (doc_idx, encoded) in doc_tokens.iter().enumerate() {
276        let mut touched: Vec<u32> =
277            encoded.iter().map(|ev| ev.centroid_id).collect();
278        touched.sort_unstable();
279        touched.dedup();
280        for cid in touched {
281            lists[cid as usize].push(doc_idx as u32);
282        }
283    }
284    InvertedFile { lists }
285}
286
287#[cfg(test)]
288mod tests {
289    use super::*;
290    use crate::distance::squared_l2;
291
292    /// Build a tiny 2-D corpus with two clear clusters of tokens.
293    fn small_corpus() -> Vec<DocumentTokens> {
294        vec![
295            DocumentTokens {
296                doc_id: 1,
297                tokens: vec![0.0, 0.0, 0.1, 0.2, -0.1, 0.1],
298                n_tokens: 3,
299            },
300            DocumentTokens {
301                doc_id: 2,
302                tokens: vec![10.0, 10.0, 10.2, 9.9, 9.8, 10.1],
303                n_tokens: 3,
304            },
305            DocumentTokens {
306                doc_id: 3,
307                tokens: vec![0.3, -0.2, 9.7, 10.2],
308                n_tokens: 2,
309            },
310        ]
311    }
312
313    fn default_params() -> IndexParams {
314        IndexParams {
315            dim: 2,
316            nbits: 2,
317            k_centroids: 2,
318            max_kmeans_iters: 50,
319        }
320    }
321
322    #[test]
323    fn build_index_encodes_every_token() {
324        let docs = small_corpus();
325        let params = default_params();
326        let expected_total: usize = docs.iter().map(|d| d.n_tokens).sum();
327
328        let index = build_index(&docs, params).unwrap();
329
330        assert_eq!(index.num_documents(), docs.len());
331        assert_eq!(index.num_tokens(), expected_total);
332        for (encoded, doc) in index.doc_tokens.iter().zip(docs.iter()) {
333            assert_eq!(encoded.len(), doc.n_tokens);
334        }
335    }
336
337    #[test]
338    fn build_index_assigns_tokens_in_each_cluster_to_the_closest_centroid() {
339        let docs = small_corpus();
340        let params = default_params();
341        let index = build_index(&docs, params).unwrap();
342
343        // The two tight clusters around (0,0) and (10,10) should produce
344        // centroids close to those means.
345        let c0 = &index.codec.centroids[0..2];
346        let c1 = &index.codec.centroids[2..4];
347
348        let (near_origin, near_ten) =
349            if squared_l2(c0, &[0.0, 0.0]) < squared_l2(c0, &[10.0, 10.0]) {
350                (c0, c1)
351            } else {
352                (c1, c0)
353            };
354
355        assert!(squared_l2(near_origin, &[0.0, 0.0]) < 1.0);
356        assert!(squared_l2(near_ten, &[10.0, 10.0]) < 1.0);
357    }
358
359    #[test]
360    fn build_index_round_trip_reconstruction_error_is_bounded() {
361        let docs = small_corpus();
362        let params = default_params();
363        let index = build_index(&docs, params).unwrap();
364
365        for (doc, encoded_doc) in docs.iter().zip(index.doc_tokens.iter()) {
366            for (token, encoded) in
367                doc.tokens.chunks_exact(params.dim).zip(encoded_doc.iter())
368            {
369                let decoded = index.codec.decode_vector(encoded).unwrap();
370                let err = squared_l2(token, &decoded).sqrt();
371                // Residuals on a 2-D toy corpus with tight clusters stay
372                // small; each bucket should cover well under 0.5 per dim.
373                assert!(
374                    err < 0.6,
375                    "reconstruction error {err} too large for token {token:?}"
376                );
377            }
378        }
379    }
380
381    #[test]
382    fn build_index_preserves_document_id_order() {
383        let docs = small_corpus();
384        let index = build_index(&docs, default_params()).unwrap();
385        let expected_ids: Vec<u64> = docs.iter().map(|d| d.doc_id).collect();
386        assert_eq!(index.doc_ids, expected_ids);
387        assert_eq!(index.position_of(2), Some(1));
388        assert_eq!(index.position_of(999), None);
389    }
390
391    #[test]
392    fn build_index_handles_document_with_no_tokens() {
393        // Empty documents are still indexable: they contribute no tokens
394        // to training but keep their slot so callers can look them up.
395        let mut docs = small_corpus();
396        docs.push(DocumentTokens {
397            doc_id: 42,
398            tokens: vec![],
399            n_tokens: 0,
400        });
401        let index = build_index(&docs, default_params()).unwrap();
402        assert_eq!(index.num_documents(), 4);
403        assert_eq!(index.doc_tokens[3].len(), 0);
404    }
405
406    #[test]
407    #[should_panic(expected = "declared")]
408    fn build_index_panics_on_mismatched_token_count() {
409        let docs = vec![DocumentTokens {
410            doc_id: 1,
411            tokens: vec![0.0, 0.0, 1.0],
412            n_tokens: 2, // says 2 but only 3 f32s and dim=2
413        }];
414        let _ = build_index(&docs, default_params()).unwrap();
415    }
416
417    #[test]
418    fn build_index_ivf_has_one_list_per_centroid() {
419        let docs = small_corpus();
420        let index = build_index(&docs, default_params()).unwrap();
421
422        assert_eq!(
423            index.ivf.num_centroids(),
424            default_params().k_centroids,
425            "IVF has one list per centroid",
426        );
427    }
428
429    #[test]
430    fn build_index_ivf_postings_cover_every_doc_that_touches_each_centroid() {
431        // Every (doc, token) pair implies the doc_idx must appear in
432        // that centroid's posting list. This is the PLAID "centroid →
433        // unique passage ids" contract: postings are indexed by doc.
434        let docs = small_corpus();
435        let index = build_index(&docs, default_params()).unwrap();
436
437        for (doc_idx, encoded_doc) in index.doc_tokens.iter().enumerate() {
438            for ev in encoded_doc {
439                let postings =
440                    index.ivf.docs_for_centroid(ev.centroid_id as usize);
441                assert!(
442                    postings.contains(&(doc_idx as u32)),
443                    "doc_idx={doc_idx} missing from centroid {} postings",
444                    ev.centroid_id,
445                );
446            }
447        }
448    }
449
450    #[test]
451    fn build_index_ivf_postings_are_unique_per_centroid() {
452        // PLAID stores centroid → unique doc ids, not token refs. A doc
453        // with multiple tokens in the same centroid must appear at most
454        // once in that centroid's list.
455        let docs = small_corpus();
456        let index = build_index(&docs, default_params()).unwrap();
457
458        for c in 0..index.ivf.num_centroids() {
459            let postings = index.ivf.docs_for_centroid(c);
460            let mut unique: Vec<u32> = postings.to_vec();
461            unique.sort_unstable();
462            unique.dedup();
463            assert_eq!(
464                unique.len(),
465                postings.len(),
466                "centroid {c} has duplicate doc entries: {postings:?}",
467            );
468        }
469    }
470
471    #[test]
472    fn build_index_dedupes_repeated_tokens_in_same_centroid() {
473        // Doc 1 has three tokens that all cluster to the same coarse
474        // centroid. The doc_idx should show up once in that centroid's
475        // posting, not three times.
476        let docs = vec![
477            DocumentTokens {
478                doc_id: 1,
479                tokens: vec![0.0, 0.0, 0.05, -0.02, -0.03, 0.01],
480                n_tokens: 3,
481            },
482            DocumentTokens {
483                doc_id: 2,
484                tokens: vec![10.0, 10.0, 10.1, 9.9],
485                n_tokens: 2,
486            },
487        ];
488        let index = build_index(&docs, default_params()).unwrap();
489
490        for c in 0..index.ivf.num_centroids() {
491            let postings = index.ivf.docs_for_centroid(c);
492            let count_of_doc_0 = postings.iter().filter(|&&d| d == 0).count();
493            assert!(
494                count_of_doc_0 <= 1,
495                "doc 0 appears {count_of_doc_0} times in centroid {c}",
496            );
497        }
498    }
499
500    #[test]
501    fn inverted_file_out_of_range_returns_empty_slice() {
502        let ivf = InvertedFile {
503            lists: vec![vec![0u32]],
504        };
505        assert_eq!(ivf.docs_for_centroid(0).len(), 1);
506        assert!(ivf.docs_for_centroid(999).is_empty());
507    }
508
509    #[test]
510    #[should_panic(expected = "at least")]
511    fn build_index_panics_when_too_few_tokens_for_k() {
512        let docs = vec![DocumentTokens {
513            doc_id: 1,
514            tokens: vec![0.0, 1.0],
515            n_tokens: 1,
516        }];
517        let params = IndexParams {
518            dim: 2,
519            nbits: 2,
520            k_centroids: 4,
521            max_kmeans_iters: 10,
522        };
523        let _ = build_index(&docs, params).unwrap();
524    }
525}