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