Skip to main content

docbert_plaid/
update.rs

1//! Incremental updates: mutate an existing [`Index`] without retraining.
2//!
3//! A freshly built index "freezes" its codec — the centroids and
4//! residual cutoffs/weights were learned from the token distribution
5//! at that point in time. In docbert's sync loop the vast majority of
6//! documents are unchanged from one sync to the next; only a handful
7//! are added, updated, or deleted. Retraining the codec every sync is
8//! wasteful: k-means and quantizer training dominate [`build_index`].
9//!
10//! [`apply_update`] takes a small mutation plan and produces a new
11//! index with deleted documents removed, upserted documents re-encoded
12//! against the *existing* codec, and the centroid → tokens inverted
13//! file rebuilt from the final encoded token list. No k-means, no
14//! quantizer retraining.
15//!
16//! Callers are expected to trigger a full rebuild (via
17//! [`build_index`]) periodically to combat codec drift as the corpus
18//! evolves — the existing codec only stays well-calibrated while the
19//! underlying token distribution doesn't change dramatically.
20//!
21//! [`build_index`]: crate::index::build_index
22
23use std::collections::HashSet;
24
25use crate::{
26    Result,
27    codec::EncodedVector,
28    index::{
29        DocumentTokens,
30        Index,
31        InvertedFile,
32        build_inverted_file_from_flat,
33    },
34};
35
36/// Mutation plan consumed by [`apply_update`].
37///
38/// Deletions are applied before upserts, so upserting a doc_id that
39/// is simultaneously listed in `deletions` is equivalent to upserting
40/// it alone. Duplicate doc_ids *within* `upserts` are rejected —
41/// the caller should pre-merge them since "last one wins" would be
42/// surprising and "first one wins" would be ambiguous.
43#[derive(Debug, Clone, Copy)]
44pub struct IndexUpdate<'a> {
45    /// Doc IDs to drop from the index.
46    pub deletions: &'a [u64],
47    /// Documents to add or replace. A document with a doc_id already
48    /// in the index has its old tokens removed before the new ones
49    /// are encoded.
50    pub upserts: &'a [DocumentTokens],
51}
52
53/// Produce a new [`Index`] reflecting `update` applied to `index`,
54/// reusing the existing codec and centroids.
55///
56/// # Errors
57///
58/// Propagates [`PlaidError::Tensor`] from the batched
59/// nearest-centroid lookup used to re-encode upserted documents.
60///
61/// # Panics
62///
63/// Panics if any upserted document's `tokens.len()` disagrees with
64/// `n_tokens * index.params.dim`, or if `upserts` contains two
65/// entries with the same `doc_id`.
66///
67/// [`PlaidError::Tensor`]: crate::PlaidError::Tensor
68pub fn apply_update(index: Index, update: IndexUpdate<'_>) -> Result<Index> {
69    let params = index.params;
70
71    for doc in update.upserts {
72        assert!(
73            doc.tokens.len() == doc.n_tokens * params.dim,
74            "apply_update: doc {} declared {} tokens but carries {} f32s (dim={})",
75            doc.doc_id,
76            doc.n_tokens,
77            doc.tokens.len(),
78            params.dim,
79        );
80    }
81
82    // Guard against ambiguous duplicate upserts early so callers get a
83    // loud failure rather than silently losing one of their writes.
84    let mut seen: HashSet<u64> = HashSet::with_capacity(update.upserts.len());
85    for doc in update.upserts {
86        assert!(
87            seen.insert(doc.doc_id),
88            "apply_update: duplicate doc_id {} in upserts",
89            doc.doc_id,
90        );
91    }
92
93    // Everything in `deletions` leaves the index. Anything in
94    // `upserts` also leaves first (so the upsert replaces the old
95    // copy when both live in the old index).
96    let mut to_remove: HashSet<u64> =
97        update.deletions.iter().copied().collect();
98    for doc in update.upserts {
99        to_remove.insert(doc.doc_id);
100    }
101
102    // Pull the existing state out. Re-materialise per-doc
103    // EncodedVector lists once so the merge/upsert loop below can
104    // think in per-document terms; we'll reflatten at the end.
105    let codec = index.codec.clone();
106    let packed_bytes = codec.packed_bytes();
107    let existing_doc_ids = index.doc_ids.clone();
108    let existing_doc_tokens: Vec<Vec<EncodedVector>> = (0..existing_doc_ids
109        .len())
110        .map(|i| index.doc_tokens_vec(i))
111        .collect();
112    drop(index);
113
114    // Retain existing docs that survive the mutation.
115    let mut new_doc_ids: Vec<u64> = Vec::with_capacity(existing_doc_ids.len());
116    let mut new_doc_tokens: Vec<Vec<EncodedVector>> =
117        Vec::with_capacity(existing_doc_tokens.len());
118    for (id, tokens) in existing_doc_ids.into_iter().zip(existing_doc_tokens) {
119        if to_remove.contains(&id) {
120            continue;
121        }
122        new_doc_ids.push(id);
123        new_doc_tokens.push(tokens);
124    }
125    let _ = packed_bytes; // consumed implicitly via codec.packed_bytes() below
126
127    // Encode every upsert's tokens against the existing codec. Doing
128    // this in one batched call lets the matmul-driven
129    // nearest-centroid lookup amortise across the whole upsert set
130    // rather than paying per-document kernel overhead.
131    let total_upsert_tokens: usize =
132        update.upserts.iter().map(|d| d.n_tokens).sum();
133    if total_upsert_tokens > 0 {
134        let mut pool: Vec<f32> =
135            Vec::with_capacity(total_upsert_tokens * params.dim);
136        for doc in update.upserts {
137            pool.extend_from_slice(&doc.tokens);
138        }
139        let (all_centroid_ids, all_codes) = codec.batch_encode_tokens(&pool)?;
140        let packed_per_token = codec.packed_bytes();
141        let mut offset = 0usize;
142        for doc in update.upserts {
143            let n = doc.n_tokens;
144            let cids = &all_centroid_ids[offset..offset + n];
145            let codes_slice = &all_codes
146                [offset * packed_per_token..(offset + n) * packed_per_token];
147            let encoded: Vec<EncodedVector> = (0..n)
148                .map(|i| EncodedVector {
149                    centroid_id: cids[i],
150                    codes: codes_slice
151                        [i * packed_per_token..(i + 1) * packed_per_token]
152                        .to_vec(),
153                })
154                .collect();
155            new_doc_ids.push(doc.doc_id);
156            new_doc_tokens.push(encoded);
157            offset += n;
158        }
159    } else {
160        // No tokens to encode, but preserve empty upserted documents'
161        // slots so callers can look them up by doc_id.
162        for doc in update.upserts {
163            new_doc_ids.push(doc.doc_id);
164            new_doc_tokens.push(Vec::new());
165        }
166    }
167
168    // Flatten into the StridedTensor-style layout and rebuild the
169    // IVF from the final token list. A full IVF rebuild is
170    // O(n_docs · avg_tokens_per_doc) and avoids tracking per-token
171    // positions inside each centroid's list.
172    //
173    // We hand the temporary Vec<Vec<EncodedVector>> to
174    // `Index::from_encoded_docs`, which flattens into the canonical
175    // layout in one pass. A temporary empty `InvertedFile` placeholder
176    // keeps the helper generic; we replace it immediately with the
177    // real IVF derived from the flat storage.
178    let mut new_index = Index::from_encoded_docs(
179        params,
180        codec,
181        new_doc_ids,
182        new_doc_tokens,
183        InvertedFile::default(),
184    );
185    new_index.ivf = build_inverted_file_from_flat(
186        &new_index.doc_centroid_ids,
187        &new_index.doc_offsets,
188        params.k_centroids,
189    );
190    Ok(new_index)
191}
192
193#[cfg(test)]
194mod tests {
195    use super::*;
196    use crate::index::{IndexParams, build_index};
197
198    fn seed_corpus() -> Vec<DocumentTokens> {
199        // Two well-separated clusters so k-means produces stable
200        // centroids we can reason about, plus a mixed doc so the IVF
201        // postings span more than one centroid per document.
202        vec![
203            DocumentTokens {
204                doc_id: 1,
205                tokens: vec![0.0, 0.0, 0.1, 0.2, -0.1, 0.1],
206                n_tokens: 3,
207            },
208            DocumentTokens {
209                doc_id: 2,
210                tokens: vec![10.0, 10.0, 10.2, 9.9, 9.8, 10.1],
211                n_tokens: 3,
212            },
213            DocumentTokens {
214                doc_id: 3,
215                tokens: vec![0.3, -0.2, 9.7, 10.2],
216                n_tokens: 2,
217            },
218        ]
219    }
220
221    fn seed_params() -> IndexParams {
222        IndexParams {
223            dim: 2,
224            nbits: 2,
225            k_centroids: 2,
226            max_kmeans_iters: 50,
227        }
228    }
229
230    fn seed_index() -> Index {
231        build_index(&seed_corpus(), seed_params()).unwrap()
232    }
233
234    /// Re-materialise every document's tokens as a
235    /// `Vec<Vec<EncodedVector>>` so the legacy assertion style
236    /// (equality checks against saved "before" snapshots) keeps
237    /// working against the flat-layout `Index`. Allocation-heavy,
238    /// only used by tests.
239    fn all_doc_tokens(index: &Index) -> Vec<Vec<EncodedVector>> {
240        (0..index.num_documents())
241            .map(|i| index.doc_tokens_vec(i))
242            .collect()
243    }
244
245    fn assert_ivf_covers_every_doc_centroid_pair(index: &Index) {
246        for doc_idx in 0..index.num_documents() {
247            for &cid in index.doc_centroid_ids(doc_idx) {
248                let list = index.ivf.docs_for_centroid(cid as usize);
249                assert!(
250                    list.contains(&(doc_idx as u32)),
251                    "missing posting for doc_idx={doc_idx} centroid={cid}",
252                );
253            }
254        }
255        // Postings are unique per centroid.
256        for c in 0..index.ivf.num_centroids() {
257            let postings = index.ivf.docs_for_centroid(c);
258            let mut sorted = postings.to_vec();
259            sorted.sort_unstable();
260            sorted.dedup();
261            assert_eq!(
262                sorted.len(),
263                postings.len(),
264                "duplicate docs in centroid {c} postings",
265            );
266        }
267    }
268
269    #[test]
270    fn apply_update_with_empty_mutations_preserves_every_document() {
271        let index = seed_index();
272        let before_ids = index.doc_ids.clone();
273        let before_tokens = all_doc_tokens(&index);
274
275        let updated = apply_update(
276            index,
277            IndexUpdate {
278                deletions: &[],
279                upserts: &[],
280            },
281        )
282        .unwrap();
283
284        assert_eq!(updated.doc_ids, before_ids);
285        assert_eq!(all_doc_tokens(&updated), before_tokens);
286        assert_ivf_covers_every_doc_centroid_pair(&updated);
287    }
288
289    #[test]
290    fn apply_update_removes_the_listed_deletions() {
291        let index = seed_index();
292        let original_tokens = index.num_tokens();
293        let removed_idx = index.position_of(2).unwrap();
294        let removed_tokens = index.doc_token_count(removed_idx);
295
296        let updated = apply_update(
297            index,
298            IndexUpdate {
299                deletions: &[2],
300                upserts: &[],
301            },
302        )
303        .unwrap();
304
305        assert_eq!(updated.doc_ids, vec![1, 3]);
306        assert_eq!(updated.num_documents(), 2);
307        assert_eq!(updated.num_tokens(), original_tokens - removed_tokens);
308        assert_ivf_covers_every_doc_centroid_pair(&updated);
309    }
310
311    #[test]
312    fn apply_update_appends_upsert_of_a_new_doc_id() {
313        let index = seed_index();
314        let new_doc = DocumentTokens {
315            doc_id: 99,
316            tokens: vec![0.05, -0.05, 0.1, 0.0],
317            n_tokens: 2,
318        };
319
320        let updated = apply_update(
321            index,
322            IndexUpdate {
323                deletions: &[],
324                upserts: std::slice::from_ref(&new_doc),
325            },
326        )
327        .unwrap();
328
329        assert_eq!(updated.doc_ids, vec![1, 2, 3, 99]);
330        let last_idx = updated.num_documents() - 1;
331        assert_eq!(updated.doc_token_count(last_idx), 2);
332        assert_ivf_covers_every_doc_centroid_pair(&updated);
333    }
334
335    #[test]
336    fn apply_update_replaces_an_existing_doc_when_upserted() {
337        let index = seed_index();
338        let replacement = DocumentTokens {
339            doc_id: 1,
340            // Different tokens in the far cluster — on replacement
341            // the encoded centroid_ids should shift accordingly.
342            tokens: vec![10.0, 10.1, 9.9, 10.0, 10.1, 9.8, 10.2, 10.0],
343            n_tokens: 4,
344        };
345        let old_encoded = {
346            let idx = index.position_of(1).unwrap();
347            index.doc_tokens_vec(idx)
348        };
349
350        let updated = apply_update(
351            index,
352            IndexUpdate {
353                deletions: &[],
354                upserts: std::slice::from_ref(&replacement),
355            },
356        )
357        .unwrap();
358
359        // doc_id 1 moves to the end because deletions happen first,
360        // then upserts are appended. Count stays the same.
361        assert_eq!(updated.doc_ids.len(), 3);
362        assert_eq!(updated.position_of(1), Some(2));
363
364        let new_encoded = updated.doc_tokens_vec(2);
365        assert_eq!(new_encoded.len(), 4);
366        assert_ne!(
367            new_encoded, old_encoded,
368            "upsert must replace the old encoded tokens",
369        );
370        assert_ivf_covers_every_doc_centroid_pair(&updated);
371    }
372
373    #[test]
374    fn apply_update_keeps_the_codec_bit_for_bit() {
375        let index = seed_index();
376        let before = index.codec.clone();
377        let upsert = DocumentTokens {
378            doc_id: 4,
379            tokens: vec![5.0, 5.0],
380            n_tokens: 1,
381        };
382
383        let updated = apply_update(
384            index,
385            IndexUpdate {
386                deletions: &[2],
387                upserts: std::slice::from_ref(&upsert),
388            },
389        )
390        .unwrap();
391
392        // Incremental updates must never touch the trained codec.
393        assert_eq!(updated.codec.centroids, before.centroids);
394        assert_eq!(updated.codec.bucket_cutoffs, before.bucket_cutoffs);
395        assert_eq!(updated.codec.bucket_weights, before.bucket_weights);
396        assert_eq!(updated.codec.nbits, before.nbits);
397        assert_eq!(updated.codec.dim, before.dim);
398    }
399
400    #[test]
401    fn apply_update_preserves_surviving_documents_verbatim() {
402        let index = seed_index();
403        // Snapshot every surviving doc's encoded tokens BEFORE the
404        // update so we can assert they're carried through unchanged.
405        let keep_ids: Vec<u64> = index
406            .doc_ids
407            .iter()
408            .copied()
409            .filter(|id| *id != 2)
410            .collect();
411        let keep_tokens: Vec<Vec<EncodedVector>> = index
412            .doc_ids
413            .iter()
414            .enumerate()
415            .filter(|(_, id)| **id != 2)
416            .map(|(i, _)| index.doc_tokens_vec(i))
417            .collect();
418
419        let updated = apply_update(
420            index,
421            IndexUpdate {
422                deletions: &[2],
423                upserts: &[],
424            },
425        )
426        .unwrap();
427
428        assert_eq!(updated.doc_ids, keep_ids);
429        assert_eq!(all_doc_tokens(&updated), keep_tokens);
430    }
431
432    #[test]
433    fn apply_update_handles_upsert_of_an_empty_document() {
434        let index = seed_index();
435        let empty = DocumentTokens {
436            doc_id: 77,
437            tokens: vec![],
438            n_tokens: 0,
439        };
440
441        let updated = apply_update(
442            index,
443            IndexUpdate {
444                deletions: &[],
445                upserts: std::slice::from_ref(&empty),
446            },
447        )
448        .unwrap();
449
450        assert!(updated.doc_ids.contains(&77));
451        assert_eq!(
452            updated.doc_token_count(updated.position_of(77).unwrap()),
453            0
454        );
455        assert_ivf_covers_every_doc_centroid_pair(&updated);
456    }
457
458    #[test]
459    #[should_panic(expected = "carries")]
460    fn apply_update_panics_on_dim_mismatch_in_upsert() {
461        let index = seed_index();
462        let bad = DocumentTokens {
463            doc_id: 5,
464            tokens: vec![1.0, 2.0, 3.0], // dim=2 so 3 f32s is invalid
465            n_tokens: 2,
466        };
467        let _ = apply_update(
468            index,
469            IndexUpdate {
470                deletions: &[],
471                upserts: std::slice::from_ref(&bad),
472            },
473        )
474        .unwrap();
475    }
476
477    #[test]
478    #[should_panic(expected = "duplicate doc_id")]
479    fn apply_update_panics_on_duplicate_upsert_doc_ids() {
480        let index = seed_index();
481        let doc_a = DocumentTokens {
482            doc_id: 1,
483            tokens: vec![0.0, 0.0],
484            n_tokens: 1,
485        };
486        let doc_b = DocumentTokens {
487            doc_id: 1,
488            tokens: vec![1.0, 1.0],
489            n_tokens: 1,
490        };
491        let _ = apply_update(
492            index,
493            IndexUpdate {
494                deletions: &[],
495                upserts: &[doc_a, doc_b],
496            },
497        )
498        .unwrap();
499    }
500}