Skip to main content

docbert_plaid/
persistence.rs

1//! Save and load a built [`Index`] to/from disk.
2//!
3//! The on-disk format is a single little-endian binary file with a small
4//! header followed by a sequence of plain f32/u32/u64 blobs. The layout
5//! is deliberately boring — no compression, no extra framing — so that
6//! every section can be read or written as one bulk copy.
7//!
8//! Layout (version 3), every field little-endian:
9//!
10//! ```text
11//! magic           : 8 bytes, b"PLAIDIDX"
12//! version         : u32     (currently 3)
13//! dim             : u32
14//! nbits           : u32
15//! k_centroids    : u32
16//! max_kmeans_iters: u32
17//! n_documents     : u64
18//! centroids       : dim * k_centroids     f32
19//! bucket_cutoffs  : 2^nbits - 1            f32
20//! bucket_weights  : 2^nbits                f32
21//! doc_ids         : n_documents            u64
22//! token_counts    : n_documents            u32  (tokens per document)
23//! centroid_ids    : total_tokens           u32
24//! residuals       : total_tokens * (dim * nbits / 8) u8
25//! ```
26//!
27//! The inverted file is *not* persisted — it's derivable from the
28//! encoded tokens in O(n_tokens) time on load and storing it would just
29//! duplicate state that's already on disk elsewhere.
30
31use std::{
32    fs::File,
33    io::{BufReader, BufWriter, Read, Write},
34    path::Path,
35};
36
37use crate::{
38    PlaidError,
39    Result,
40    codec::{ResidualCodec, packed_bytes_per_vector},
41    index::{Index, IndexParams, build_inverted_file_from_flat},
42};
43
44const MAGIC: &[u8; 8] = b"PLAIDIDX";
45/// Format versions in the wild:
46///
47/// - `1`: unpacked residual codes, one byte per residual dimension.
48/// - `2`: LSB-first bit-packed codes at `nbits ∈ {1, 2, 4, 8}`;
49///   `(dim * nbits) / 8` bytes per token, stored as per-token
50///   interleaved `(centroid_id, residual)` records.
51/// - `3`: same codes, but centroid ids and residuals are two
52///   contiguous sections so loading a multi-GB index is a pair of
53///   bulk copies instead of a per-token de-interleave.
54///
55/// Only version 3 is readable; older files must be regenerated from
56/// the stored embeddings (`docbert rebuild`).
57const FORMAT_VERSION: u32 = 3;
58
59/// Write `index` to `path`, creating or truncating the file as needed.
60///
61/// # Errors
62///
63/// Returns [`PlaidError::Io`] for any underlying I/O failure, or
64/// [`PlaidError::InvalidIndex`] if an encoded token's `codes` buffer
65/// doesn't match the codec's advertised packed length (a caller bug
66/// worth flagging loudly on write rather than producing a file that
67/// would fail to load back).
68///
69/// [`PlaidError::Io`]: crate::PlaidError::Io
70/// [`PlaidError::InvalidIndex`]: crate::PlaidError::InvalidIndex
71pub fn save(index: &Index, path: &Path) -> Result<()> {
72    let file = File::create(path)?;
73    let mut writer = BufWriter::new(file);
74    write_index(index, &mut writer)?;
75    writer.flush()?;
76    Ok(())
77}
78
79/// Read an [`Index`] back from `path`.
80///
81/// # Errors
82///
83/// Returns [`PlaidError::Io`] for any read failure, or
84/// [`PlaidError::InvalidIndex`] if the file's magic, version, or
85/// header fields don't match this crate's format.
86///
87/// [`PlaidError::Io`]: crate::PlaidError::Io
88/// [`PlaidError::InvalidIndex`]: crate::PlaidError::InvalidIndex
89pub fn load(path: &Path) -> Result<Index> {
90    let file = File::open(path)?;
91    let mut reader = BufReader::new(file);
92    read_index(&mut reader)
93}
94
95fn write_index<W: Write>(index: &Index, w: &mut W) -> Result<()> {
96    w.write_all(MAGIC)?;
97    write_u32(w, FORMAT_VERSION)?;
98
99    let params = &index.params;
100    write_u32(w, params.dim as u32)?;
101    write_u32(w, params.nbits)?;
102    write_u32(w, params.k_centroids as u32)?;
103    write_u32(w, params.max_kmeans_iters as u32)?;
104    write_u64(w, index.doc_ids.len() as u64)?;
105
106    write_f32_slice(w, &index.codec.centroids)?;
107    write_f32_slice(w, &index.codec.bucket_cutoffs)?;
108    write_f32_slice(w, &index.codec.bucket_weights)?;
109
110    write_u64_slice(w, &index.doc_ids)?;
111
112    let token_counts: Vec<u32> = (0..index.num_documents())
113        .map(|i| index.doc_token_count(i) as u32)
114        .collect();
115    write_u32_slice(w, &token_counts)?;
116
117    let packed_bytes = packed_bytes_per_vector(params.dim, params.nbits);
118    if index.doc_residual_bytes.len() != index.num_tokens() * packed_bytes {
119        return Err(PlaidError::InvalidIndex(format!(
120            "doc_residual_bytes length {} is not tokens ({}) × packed_bytes ({packed_bytes})",
121            index.doc_residual_bytes.len(),
122            index.num_tokens(),
123        )));
124    }
125
126    write_u32_slice(w, &index.doc_centroid_ids)?;
127    w.write_all(&index.doc_residual_bytes)?;
128
129    Ok(())
130}
131
132fn read_index<R: Read>(r: &mut R) -> Result<Index> {
133    let mut magic = [0u8; 8];
134    r.read_exact(&mut magic)?;
135    if &magic != MAGIC {
136        return Err(PlaidError::InvalidIndex(
137            "not a docbert-plaid index (magic bytes mismatch)".into(),
138        ));
139    }
140
141    let version = read_u32(r)?;
142    if version != FORMAT_VERSION {
143        return Err(PlaidError::InvalidIndex(format!(
144            "unsupported plaid index version {version}, expected \
145             {FORMAT_VERSION}; run `docbert rebuild` to regenerate the \
146             index from the stored embeddings",
147        )));
148    }
149
150    let dim = read_u32(r)? as usize;
151    let nbits = read_u32(r)?;
152    let k_centroids = read_u32(r)? as usize;
153    let max_kmeans_iters = read_u32(r)? as usize;
154    let n_documents = read_u64(r)? as usize;
155
156    if dim == 0 || k_centroids == 0 || !matches!(nbits, 1 | 2 | 4 | 8) {
157        return Err(PlaidError::InvalidIndex(
158            "plaid index header has invalid dim/k_centroids/nbits".into(),
159        ));
160    }
161
162    let params = IndexParams {
163        dim,
164        nbits,
165        k_centroids,
166        max_kmeans_iters,
167    };
168
169    let centroids = read_f32_vec(r, k_centroids * dim)?;
170    let num_buckets = 1usize << nbits;
171    let bucket_cutoffs = read_f32_vec(r, num_buckets - 1)?;
172    let bucket_weights = read_f32_vec(r, num_buckets)?;
173
174    let doc_ids = read_u64_vec(r, n_documents)?;
175    let token_counts = read_u32_vec(r, n_documents)?;
176    let packed_bytes = packed_bytes_per_vector(dim, nbits);
177
178    let total_tokens: usize = token_counts.iter().map(|&c| c as usize).sum();
179
180    // Contiguous sections: each is one bulk copy into place.
181    let doc_centroid_ids = read_u32_vec(r, total_tokens)?;
182    let doc_residual_bytes = read_u8_vec(r, total_tokens * packed_bytes)?;
183
184    // Validate ids after the copy loops: a whole-slice scan
185    // vectorizes, a branch per token does not.
186    if let Some(&bad) = doc_centroid_ids
187        .iter()
188        .find(|&&id| (id as usize) >= k_centroids)
189    {
190        return Err(PlaidError::InvalidIndex(format!(
191            "centroid_id {bad} out of range 0..{k_centroids}",
192        )));
193    }
194
195    let mut doc_offsets: Vec<usize> = Vec::with_capacity(n_documents + 1);
196    doc_offsets.push(0);
197    let mut acc = 0usize;
198    for &count in &token_counts {
199        acc += count as usize;
200        doc_offsets.push(acc);
201    }
202
203    let ivf = build_inverted_file_from_flat(
204        &doc_centroid_ids,
205        &doc_offsets,
206        k_centroids,
207    );
208
209    let codec = ResidualCodec {
210        nbits,
211        dim,
212        centroids,
213        bucket_cutoffs,
214        bucket_weights,
215    };
216    codec.validate()?;
217
218    Ok(Index {
219        params,
220        codec,
221        doc_ids,
222        doc_centroid_ids,
223        doc_residual_bytes,
224        doc_offsets,
225        ivf,
226    })
227}
228
229fn write_u32<W: Write>(w: &mut W, v: u32) -> Result<()> {
230    w.write_all(&v.to_le_bytes())?;
231    Ok(())
232}
233
234fn write_u64<W: Write>(w: &mut W, v: u64) -> Result<()> {
235    w.write_all(&v.to_le_bytes())?;
236    Ok(())
237}
238
239fn write_f32_slice<W: Write>(w: &mut W, slice: &[f32]) -> Result<()> {
240    w.write_all(bytemuck::cast_slice(slice))?;
241    Ok(())
242}
243
244fn write_u32_slice<W: Write>(w: &mut W, slice: &[u32]) -> Result<()> {
245    w.write_all(bytemuck::cast_slice(slice))?;
246    Ok(())
247}
248
249fn write_u64_slice<W: Write>(w: &mut W, slice: &[u64]) -> Result<()> {
250    w.write_all(bytemuck::cast_slice(slice))?;
251    Ok(())
252}
253
254fn read_u32<R: Read>(r: &mut R) -> Result<u32> {
255    let mut buf = [0u8; 4];
256    r.read_exact(&mut buf)?;
257    Ok(u32::from_le_bytes(buf))
258}
259
260fn read_u64<R: Read>(r: &mut R) -> Result<u64> {
261    let mut buf = [0u8; 8];
262    r.read_exact(&mut buf)?;
263    Ok(u64::from_le_bytes(buf))
264}
265
266fn read_f32_vec<R: Read>(r: &mut R, n: usize) -> Result<Vec<f32>> {
267    let mut out = vec![0.0f32; n];
268    r.read_exact(bytemuck::cast_slice_mut(&mut out))?;
269    Ok(out)
270}
271
272fn read_u8_vec<R: Read>(r: &mut R, n: usize) -> Result<Vec<u8>> {
273    let mut out = vec![0u8; n];
274    r.read_exact(&mut out)?;
275    Ok(out)
276}
277
278fn read_u32_vec<R: Read>(r: &mut R, n: usize) -> Result<Vec<u32>> {
279    let mut out = vec![0u32; n];
280    r.read_exact(bytemuck::cast_slice_mut(&mut out))?;
281    Ok(out)
282}
283
284fn read_u64_vec<R: Read>(r: &mut R, n: usize) -> Result<Vec<u64>> {
285    let mut out = vec![0u64; n];
286    r.read_exact(bytemuck::cast_slice_mut(&mut out))?;
287    Ok(out)
288}
289
290#[cfg(test)]
291mod tests {
292    use super::*;
293    use crate::index::{DocumentTokens, build_index};
294
295    fn small_corpus() -> Vec<DocumentTokens> {
296        vec![
297            DocumentTokens {
298                doc_id: 10,
299                tokens: vec![0.0, 0.0, 0.1, 0.2, -0.1, 0.1],
300                n_tokens: 3,
301            },
302            DocumentTokens {
303                doc_id: 20,
304                tokens: vec![10.0, 10.0, 10.2, 9.9, 9.8, 10.1],
305                n_tokens: 3,
306            },
307            DocumentTokens {
308                doc_id: 30,
309                tokens: vec![0.3, -0.2, 9.7, 10.2],
310                n_tokens: 2,
311            },
312        ]
313    }
314
315    fn default_params() -> IndexParams {
316        IndexParams {
317            dim: 2,
318            nbits: 2,
319            k_centroids: 2,
320            max_kmeans_iters: 50,
321        }
322    }
323
324    #[test]
325    fn round_trip_preserves_codec_parameters_and_doc_ids() {
326        let tmp = tempfile::tempdir().unwrap();
327        let path = tmp.path().join("index.plaid");
328        let index = build_index(&small_corpus(), default_params()).unwrap();
329
330        save(&index, &path).unwrap();
331        let loaded = load(&path).unwrap();
332
333        assert_eq!(loaded.params.dim, index.params.dim);
334        assert_eq!(loaded.params.nbits, index.params.nbits);
335        assert_eq!(loaded.params.k_centroids, index.params.k_centroids);
336        assert_eq!(
337            loaded.params.max_kmeans_iters,
338            index.params.max_kmeans_iters
339        );
340        assert_eq!(loaded.doc_ids, index.doc_ids);
341    }
342
343    #[test]
344    fn round_trip_preserves_codec_tables_byte_for_byte() {
345        let tmp = tempfile::tempdir().unwrap();
346        let path = tmp.path().join("index.plaid");
347        let index = build_index(&small_corpus(), default_params()).unwrap();
348
349        save(&index, &path).unwrap();
350        let loaded = load(&path).unwrap();
351
352        assert_eq!(loaded.codec.centroids, index.codec.centroids);
353        assert_eq!(loaded.codec.bucket_cutoffs, index.codec.bucket_cutoffs);
354        assert_eq!(loaded.codec.bucket_weights, index.codec.bucket_weights);
355    }
356
357    #[test]
358    fn round_trip_preserves_encoded_tokens() {
359        let tmp = tempfile::tempdir().unwrap();
360        let path = tmp.path().join("index.plaid");
361        let index = build_index(&small_corpus(), default_params()).unwrap();
362
363        save(&index, &path).unwrap();
364        let loaded = load(&path).unwrap();
365
366        assert_eq!(loaded.num_documents(), index.num_documents());
367        for i in 0..index.num_documents() {
368            assert_eq!(loaded.doc_tokens_vec(i), index.doc_tokens_vec(i));
369        }
370    }
371
372    #[test]
373    fn round_trip_rebuilds_the_inverted_file() {
374        let tmp = tempfile::tempdir().unwrap();
375        let path = tmp.path().join("index.plaid");
376        let index = build_index(&small_corpus(), default_params()).unwrap();
377
378        save(&index, &path).unwrap();
379        let loaded = load(&path).unwrap();
380
381        assert_eq!(loaded.ivf.num_centroids(), index.ivf.num_centroids());
382        assert_eq!(
383            loaded.ivf.total_doc_postings(),
384            index.ivf.total_doc_postings(),
385        );
386        for c in 0..index.ivf.num_centroids() {
387            let want = index.ivf.docs_for_centroid(c);
388            let got = loaded.ivf.docs_for_centroid(c);
389            assert_eq!(got, want, "IVF list for centroid {c} differs");
390        }
391    }
392
393    #[test]
394    fn round_trip_preserves_search_results_exactly() {
395        let tmp = tempfile::tempdir().unwrap();
396        let path = tmp.path().join("index.plaid");
397        let index = build_index(&small_corpus(), default_params()).unwrap();
398        save(&index, &path).unwrap();
399        let loaded = load(&path).unwrap();
400
401        let query = [0.05f32, 0.1, 9.9, 10.1];
402        let params = crate::search::SearchParams {
403            top_k: 3,
404            n_probe: 2,
405            n_candidate_docs: None,
406            centroid_score_threshold: None,
407        };
408
409        let a = crate::search::search(&index, &query, params).unwrap();
410        let b = crate::search::search(&loaded, &query, params).unwrap();
411        assert_eq!(a, b);
412    }
413
414    #[test]
415    fn round_trip_handles_empty_documents() {
416        let tmp = tempfile::tempdir().unwrap();
417        let path = tmp.path().join("index.plaid");
418        let mut docs = small_corpus();
419        docs.push(DocumentTokens {
420            doc_id: 999,
421            tokens: vec![],
422            n_tokens: 0,
423        });
424        let index = build_index(&docs, default_params()).unwrap();
425
426        save(&index, &path).unwrap();
427        let loaded = load(&path).unwrap();
428
429        assert_eq!(loaded.doc_ids, index.doc_ids);
430        assert_eq!(loaded.doc_token_count(loaded.num_documents() - 1), 0);
431    }
432
433    #[test]
434    fn load_rejects_files_with_wrong_magic() {
435        let tmp = tempfile::tempdir().unwrap();
436        let path = tmp.path().join("bogus.plaid");
437        std::fs::write(&path, b"NOTPLAID").unwrap();
438
439        let err = load(&path).unwrap_err();
440        assert!(
441            matches!(err, PlaidError::InvalidIndex(ref m) if m.contains("magic")),
442            "expected InvalidIndex with magic message, got {err:?}",
443        );
444    }
445
446    #[test]
447    fn load_rejects_unknown_format_version() {
448        let tmp = tempfile::tempdir().unwrap();
449        let path = tmp.path().join("future.plaid");
450        let mut buf = Vec::new();
451        buf.extend_from_slice(MAGIC);
452        buf.extend_from_slice(&999u32.to_le_bytes());
453        std::fs::write(&path, &buf).unwrap();
454
455        let err = load(&path).unwrap_err();
456        assert!(
457            matches!(err, PlaidError::InvalidIndex(ref m) if m.contains("version")),
458            "expected InvalidIndex with version message, got {err:?}",
459        );
460    }
461
462    #[test]
463    fn load_rejects_legacy_versions_with_rebuild_instructions() {
464        // Pre-1.0 formats (v1 unpacked, v2 interleaved) must be
465        // regenerated; the loader rejects them and says how.
466        let tmp = tempfile::tempdir().unwrap();
467        for version in [1u32, 2] {
468            let path = tmp.path().join(format!("legacy_v{version}.plaid"));
469            let mut buf = Vec::new();
470            buf.extend_from_slice(MAGIC);
471            buf.extend_from_slice(&version.to_le_bytes());
472            std::fs::write(&path, &buf).unwrap();
473
474            let err = load(&path).unwrap_err();
475            let PlaidError::InvalidIndex(msg) = &err else {
476                panic!("expected InvalidIndex for v{version}, got {err:?}");
477            };
478            assert!(msg.contains("version"), "v{version}: {msg}");
479            assert!(msg.contains("docbert rebuild"), "v{version}: {msg}");
480        }
481    }
482}