Skip to main content

docbert_plaid/
codec.rs

1//! Residual quantization codec for PLAID.
2//!
3//! Once k-means has produced a set of coarse centroids, every token
4//! embedding can be represented as:
5//!
6//! ```text
7//! token ≈ centroid[centroid_id] + decode(residual_codes)
8//! ```
9//!
10//! The residual is the element-wise difference between the token and its
11//! nearest centroid. Each residual dimension is then placed into one of
12//! `2^nbits` buckets according to a precomputed set of cutoffs, and the
13//! bucket index (0…2ⁿ-1) is what we store on disk. At read time the
14//! bucket index is mapped back to a reconstruction value via
15//! `bucket_weights` and added to the centroid, yielding an approximate
16//! copy of the original token.
17//!
18//! This module exposes the codec state and the encode/decode operations
19//! assuming the codec has already been trained. Bucket cutoffs and
20//! weights are learned from a sample of residuals via
21//! [`train_quantizer`].
22//!
23//! Storage layout: residual codes are LSB-first bit-packed at `nbits`
24//! bits each. Supported widths are `{1, 2, 4, 8}` — enough to cover
25//! every value the ColBERTv2/PLAID papers use in practice. For a 128-d
26//! embedding at 2 bits, this is 32 bytes per token (vs. 128 bytes
27//! unpacked), matching the paper's §4.5 packed-index layout.
28
29use crate::{
30    PlaidError,
31    Result,
32    distance::squared_l2,
33    kmeans::nearest_centroid,
34};
35
36/// A trained residual-quantization codec.
37///
38/// Cutoffs partition the real line into `2^nbits` buckets. `bucket_cutoffs`
39/// holds the `2^nbits - 1` internal boundaries in ascending order;
40/// `bucket_weights` holds the `2^nbits` reconstruction values used when
41/// decoding. Both are codec-wide: the same cutoffs/weights are applied to
42/// every residual dimension of every token.
43#[derive(Debug, Clone)]
44pub struct ResidualCodec {
45    /// Number of bits per residual dimension. Typically 2 or 4.
46    pub nbits: u32,
47    /// Dimensionality of the (original) token embeddings.
48    pub dim: usize,
49    /// Flat row-major coarse centroids, `k × dim`.
50    pub centroids: Vec<f32>,
51    /// `(2^nbits) - 1` ascending cutoff values for bucketing residuals.
52    pub bucket_cutoffs: Vec<f32>,
53    /// `2^nbits` reconstruction values, one per bucket.
54    pub bucket_weights: Vec<f32>,
55}
56
57/// A single encoded token: a centroid reference plus a bit-packed
58/// buffer of per-dim bucket codes.
59///
60/// The `codes` buffer holds `dim` quantization codes packed LSB-first
61/// at `nbits` bits each. For the supported widths of 1, 2, 4, and 8
62/// bits the buffer length is `(dim * nbits) / 8` (dim is expected to
63/// be a multiple of `8/nbits` so code positions don't span bytes —
64/// ColBERT dims are 128 or 96, which satisfies that constraint for
65/// every supported `nbits`).
66#[derive(Debug, Clone, PartialEq, Eq)]
67pub struct EncodedVector {
68    /// Index of the coarse centroid this token was quantized against.
69    pub centroid_id: u32,
70    /// Bit-packed bucket codes. Use [`ResidualCodec::read_code`] or the
71    /// codec's `decode_vector` to pull values out.
72    pub codes: Vec<u8>,
73}
74
75/// Number of bytes required to pack `dim` codes at `nbits` bits each.
76///
77/// Panics if `nbits` is not one of the supported packed widths.
78pub fn packed_bytes_per_vector(dim: usize, nbits: u32) -> usize {
79    assert_supported_nbits(nbits);
80    (dim * nbits as usize).div_ceil(8)
81}
82
83fn assert_supported_nbits(nbits: u32) {
84    assert!(
85        matches!(nbits, 1 | 2 | 4 | 8),
86        "packed codec: nbits must be 1, 2, 4, or 8 (got {nbits})",
87    );
88}
89
90/// Pack `unpacked` (one byte per code, values in `0 .. 2^nbits`) into
91/// an LSB-first bit-packed buffer.
92fn pack_codes(unpacked: &[u8], nbits: u32) -> Vec<u8> {
93    assert_supported_nbits(nbits);
94    if nbits == 8 {
95        return unpacked.to_vec();
96    }
97    let codes_per_byte = 8 / nbits as usize;
98    let mask: u8 = ((1u16 << nbits) - 1) as u8;
99    let n_bytes = unpacked.len().div_ceil(codes_per_byte);
100    let mut packed = vec![0u8; n_bytes];
101    for (i, &code) in unpacked.iter().enumerate() {
102        let byte_idx = i / codes_per_byte;
103        let bit_off = (i % codes_per_byte) * nbits as usize;
104        packed[byte_idx] |= (code & mask) << bit_off;
105    }
106    packed
107}
108
109/// Read the code at logical position `i` from a packed buffer.
110pub fn read_code(packed: &[u8], i: usize, nbits: u32) -> u8 {
111    assert_supported_nbits(nbits);
112    if nbits == 8 {
113        return packed[i];
114    }
115    let codes_per_byte = 8 / nbits as usize;
116    let mask: u8 = ((1u16 << nbits) - 1) as u8;
117    let byte_idx = i / codes_per_byte;
118    let bit_off = (i % codes_per_byte) * nbits as usize;
119    (packed[byte_idx] >> bit_off) & mask
120}
121
122/// Precomputed 256-entry lookup table mapping every possible packed
123/// byte to the sequence of `bucket_weights` values it decodes to.
124///
125/// PLAID §4.5 notes that naive decompression pays a chain of
126/// shift-mask-weight-lookup operations per residual dimension; a
127/// one-off table that already composes the shift/mask with the weight
128/// lookup reduces decoding to a single load per code position. For
129/// `nbits=2` the whole table is `256 × 4` f32 = 4 KiB and easily stays
130/// in L1.
131pub struct DecodeTable {
132    /// `weights[b * codes_per_byte + k]` = weight for the `k`-th code
133    /// position inside packed byte value `b`.
134    weights: Vec<f32>,
135    codes_per_byte: usize,
136    nbits: u32,
137}
138
139impl DecodeTable {
140    /// Build the table for `codec`. Call once per search/decode batch
141    /// and reuse across every encoded vector.
142    pub fn new(codec: &ResidualCodec) -> Self {
143        assert_supported_nbits(codec.nbits);
144        let codes_per_byte = 8 / codec.nbits as usize;
145        let entries = 256;
146        let mut weights = vec![0.0f32; entries * codes_per_byte];
147        let mask: u8 = ((1u16 << codec.nbits) - 1) as u8;
148        for b in 0u16..256 {
149            let byte = b as u8;
150            for k in 0..codes_per_byte {
151                let code = (byte >> (k * codec.nbits as usize)) & mask;
152                weights[b as usize * codes_per_byte + k] =
153                    codec.bucket_weights[code as usize];
154            }
155        }
156        Self {
157            weights,
158            codes_per_byte,
159            nbits: codec.nbits,
160        }
161    }
162
163    /// Weights for the `codes_per_byte` positions inside packed byte
164    /// `byte`. Length always equals `codes_per_byte`.
165    pub fn weights_for(&self, byte: u8) -> &[f32] {
166        let start = byte as usize * self.codes_per_byte;
167        &self.weights[start..start + self.codes_per_byte]
168    }
169
170    /// Number of codes packed into one byte at this table's `nbits`.
171    pub fn codes_per_byte(&self) -> usize {
172        self.codes_per_byte
173    }
174
175    /// Bit-width the table was built for.
176    pub fn nbits(&self) -> u32 {
177        self.nbits
178    }
179}
180
181impl ResidualCodec {
182    /// Number of buckets this codec partitions the residual space into.
183    pub fn num_buckets(&self) -> usize {
184        1usize << self.nbits
185    }
186
187    /// Number of coarse centroids stored.
188    pub fn num_centroids(&self) -> usize {
189        self.centroids.len() / self.dim
190    }
191
192    /// Number of packed bytes each encoded vector uses.
193    pub fn packed_bytes(&self) -> usize {
194        packed_bytes_per_vector(self.dim, self.nbits)
195    }
196
197    /// Validate internal shape invariants. Called automatically by
198    /// encode/decode; exposed so callers loading a codec from disk can
199    /// fail fast.
200    ///
201    /// # Errors
202    ///
203    /// Returns [`PlaidError::InvalidCodec`] with a description of the
204    /// constraint that's violated.
205    pub fn validate(&self) -> Result<()> {
206        if self.dim == 0 {
207            return Err(PlaidError::InvalidCodec(
208                "codec: dim must be positive".into(),
209            ));
210        }
211        if !matches!(self.nbits, 1 | 2 | 4 | 8) {
212            return Err(PlaidError::InvalidCodec(format!(
213                "codec: nbits must be 1, 2, 4, or 8, got {}",
214                self.nbits
215            )));
216        }
217        if !self.centroids.len().is_multiple_of(self.dim)
218            || self.centroids.is_empty()
219        {
220            return Err(PlaidError::InvalidCodec(format!(
221                "codec: centroids length {} is not a positive multiple of dim {}",
222                self.centroids.len(),
223                self.dim,
224            )));
225        }
226        let expected_buckets = self.num_buckets();
227        if self.bucket_weights.len() != expected_buckets {
228            return Err(PlaidError::InvalidCodec(format!(
229                "codec: expected {} bucket_weights, got {}",
230                expected_buckets,
231                self.bucket_weights.len(),
232            )));
233        }
234        if self.bucket_cutoffs.len() != expected_buckets - 1 {
235            return Err(PlaidError::InvalidCodec(format!(
236                "codec: expected {} bucket_cutoffs, got {}",
237                expected_buckets - 1,
238                self.bucket_cutoffs.len(),
239            )));
240        }
241        for pair in self.bucket_cutoffs.windows(2) {
242            if pair[0] > pair[1] || pair[0].is_nan() || pair[1].is_nan() {
243                return Err(PlaidError::InvalidCodec(
244                    "codec: bucket_cutoffs must be non-decreasing and finite"
245                        .into(),
246                ));
247            }
248        }
249        Ok(())
250    }
251
252    /// Encode a single token embedding.
253    ///
254    /// Finds the nearest centroid, computes the residual, and quantizes
255    /// each dimension against `bucket_cutoffs`.
256    ///
257    /// # Errors
258    ///
259    /// Returns [`PlaidError::InvalidCodec`] if this codec fails its
260    /// shape invariants.
261    ///
262    /// # Panics
263    ///
264    /// Panics if `vector.len() != dim`.
265    pub fn encode_vector(&self, vector: &[f32]) -> Result<EncodedVector> {
266        self.validate()?;
267        assert_eq!(
268            vector.len(),
269            self.dim,
270            "encode_vector: expected {} dims, got {}",
271            self.dim,
272            vector.len(),
273        );
274
275        let centroid_id = nearest_centroid(vector, &self.centroids, self.dim);
276        let centroid_slice = &self.centroids
277            [centroid_id * self.dim..(centroid_id + 1) * self.dim];
278
279        let unpacked: Vec<u8> = vector
280            .iter()
281            .zip(centroid_slice.iter())
282            .map(|(v, c)| bucket_for_value(*v - *c, &self.bucket_cutoffs))
283            .collect();
284        let codes = pack_codes(&unpacked, self.nbits);
285
286        Ok(EncodedVector {
287            centroid_id: centroid_id as u32,
288            codes,
289        })
290    }
291
292    /// Encode every token in a flat `n × dim` buffer in one batched
293    /// pass, returning the per-token centroid id and a flat `n × dim`
294    /// code buffer.
295    ///
296    /// The expensive step — the nearest-centroid lookup — runs as a
297    /// single matmul through [`crate::kmeans::assign_points`], which
298    /// uses candle's GEMM (CPU or CUDA depending on build). The
299    /// residual + bucket loop stays scalar because per-element
300    /// `searchsorted` would otherwise require either a 3-D broadcast
301    /// against the cutoffs table or a per-cutoff kernel launch — both
302    /// less efficient than a tight Rust loop over the small cutoffs
303    /// vector. Returning the codes flat avoids `n` `Vec<u8>`
304    /// allocations; callers split into per-token slices as needed.
305    ///
306    /// # Errors
307    ///
308    /// Returns [`PlaidError::InvalidCodec`] if the codec fails its
309    /// shape invariants, or [`PlaidError::Tensor`] if the
310    /// matmul-driven nearest-centroid lookup fails.
311    ///
312    /// # Panics
313    ///
314    /// Panics if `tokens.len() % dim != 0` or if `tokens` is empty.
315    ///
316    /// [`PlaidError::InvalidCodec`]: crate::PlaidError::InvalidCodec
317    /// [`PlaidError::Tensor`]: crate::PlaidError::Tensor
318    pub fn batch_encode_tokens(
319        &self,
320        tokens: &[f32],
321    ) -> Result<(Vec<u32>, Vec<u8>)> {
322        self.validate()?;
323        assert!(
324            tokens.len().is_multiple_of(self.dim),
325            "batch_encode_tokens: tokens length {} is not a multiple of dim {}",
326            tokens.len(),
327            self.dim,
328        );
329        let n = tokens.len() / self.dim;
330        if n == 0 {
331            return Ok((Vec::new(), Vec::new()));
332        }
333
334        let assignments =
335            crate::kmeans::assign_points(tokens, &self.centroids, self.dim)?;
336
337        let packed_per_token = self.packed_bytes();
338        let mut centroid_ids: Vec<u32> = Vec::with_capacity(n);
339        let mut packed_codes: Vec<u8> =
340            Vec::with_capacity(n * packed_per_token);
341        let mut scratch: Vec<u8> = Vec::with_capacity(self.dim);
342        for (token, &cluster) in
343            tokens.chunks_exact(self.dim).zip(assignments.iter())
344        {
345            let centroid_slice =
346                &self.centroids[cluster * self.dim..(cluster + 1) * self.dim];
347            scratch.clear();
348            for (t, c) in token.iter().zip(centroid_slice.iter()) {
349                scratch.push(bucket_for_value(*t - *c, &self.bucket_cutoffs));
350            }
351            packed_codes.extend(pack_codes(&scratch, self.nbits));
352            centroid_ids.push(cluster as u32);
353        }
354        Ok((centroid_ids, packed_codes))
355    }
356
357    /// Reconstruct an approximate token embedding from its codes.
358    ///
359    /// # Errors
360    ///
361    /// Returns [`PlaidError::InvalidCodec`] if this codec fails its
362    /// shape invariants.
363    ///
364    /// # Panics
365    ///
366    /// Panics if `codes.len() != dim` or if any code is out of range.
367    ///
368    /// [`PlaidError::InvalidCodec`]: crate::PlaidError::InvalidCodec
369    pub fn decode_vector(&self, encoded: &EncodedVector) -> Result<Vec<f32>> {
370        let table = DecodeTable::new(self);
371        self.decode_vector_with_table(encoded, &table)
372    }
373
374    /// Decode an encoded vector using a pre-built [`DecodeTable`].
375    ///
376    /// Callers that decode many vectors in a row (e.g., the search
377    /// path's per-candidate decode loop) should build the table once
378    /// outside the loop and reuse it here — each call then amounts to
379    /// one table load per packed byte plus the centroid add.
380    ///
381    /// # Errors
382    ///
383    /// Returns [`PlaidError::InvalidCodec`] if this codec fails its
384    /// shape invariants.
385    ///
386    /// [`PlaidError::InvalidCodec`]: crate::PlaidError::InvalidCodec
387    pub fn decode_vector_with_table(
388        &self,
389        encoded: &EncodedVector,
390        table: &DecodeTable,
391    ) -> Result<Vec<f32>> {
392        self.validate()?;
393        assert_eq!(
394            table.nbits, self.nbits,
395            "decode_vector_with_table: table nbits {} != codec nbits {}",
396            table.nbits, self.nbits,
397        );
398        let expected_bytes = self.packed_bytes();
399        assert_eq!(
400            encoded.codes.len(),
401            expected_bytes,
402            "decode_vector_with_table: expected {expected_bytes} packed bytes, got {}",
403            encoded.codes.len(),
404        );
405        let centroid_id = encoded.centroid_id as usize;
406        assert!(
407            centroid_id < self.num_centroids(),
408            "decode_vector_with_table: centroid_id {} out of range 0..{}",
409            centroid_id,
410            self.num_centroids(),
411        );
412
413        let centroid_slice = &self.centroids
414            [centroid_id * self.dim..(centroid_id + 1) * self.dim];
415        let codes_per_byte = table.codes_per_byte;
416
417        let mut out = Vec::with_capacity(self.dim);
418        for (byte_idx, &byte) in encoded.codes.iter().enumerate() {
419            let weights = table.weights_for(byte);
420            let base_dim = byte_idx * codes_per_byte;
421            for (k, &w) in weights.iter().enumerate() {
422                let dim_idx = base_dim + k;
423                if dim_idx >= self.dim {
424                    break;
425                }
426                out.push(centroid_slice[dim_idx] + w);
427            }
428        }
429        Ok(out)
430    }
431
432    /// Return the squared L2 reconstruction error for `vector` under
433    /// this codec. Useful as a lightweight codec-quality probe in tests
434    /// and evaluation scripts.
435    ///
436    /// # Errors
437    ///
438    /// Returns [`PlaidError::InvalidCodec`] if this codec fails its
439    /// shape invariants.
440    ///
441    /// [`PlaidError::InvalidCodec`]: crate::PlaidError::InvalidCodec
442    pub fn reconstruction_error(&self, vector: &[f32]) -> Result<f32> {
443        let encoded = self.encode_vector(vector)?;
444        let decoded = self.decode_vector(&encoded)?;
445        Ok(squared_l2(vector, &decoded))
446    }
447}
448
449/// Learn bucket cutoffs and reconstruction weights from a sample of
450/// residual values.
451///
452/// The returned tuple is `(bucket_cutoffs, bucket_weights)` with
453/// `2^nbits - 1` cutoffs and `2^nbits` weights, ready to plug into a
454/// [`ResidualCodec`]. Buckets are equal-quantile slices of the input:
455/// cutoffs are picked at `i / (2^nbits)` quantile positions, and each
456/// weight is the arithmetic mean of the residuals falling into that
457/// bucket. This matches fast-plaid's "fit" step and keeps the codec
458/// unbiased on the training distribution.
459///
460/// `residuals` does not need to be sorted; this function sorts a
461/// locally-owned copy. NaN values are rejected up front since they
462/// would poison the sort order.
463///
464/// # Panics
465///
466/// Panics if `residuals` is empty, if `nbits` is zero or exceeds 8, or
467/// if `residuals` contains NaN.
468pub fn train_quantizer(residuals: &[f32], nbits: u32) -> (Vec<f32>, Vec<f32>) {
469    assert!(!residuals.is_empty(), "train_quantizer: empty sample");
470    assert!(
471        nbits > 0 && nbits <= 8,
472        "train_quantizer: nbits must be in 1..=8, got {nbits}"
473    );
474    assert!(
475        residuals.iter().all(|v| !v.is_nan()),
476        "train_quantizer: residual sample contains NaN"
477    );
478
479    let num_buckets = 1usize << nbits;
480    let n = residuals.len();
481
482    let mut sorted = residuals.to_vec();
483    // We already rejected NaN above, so `total_cmp` is a strict
484    // ordering and doesn't need the `partial_cmp().unwrap()` dance.
485    sorted.sort_by(|a, b| a.total_cmp(b));
486
487    let bucket_bounds = |i: usize| -> (usize, usize) {
488        let start = i * n / num_buckets;
489        let end = if i + 1 == num_buckets {
490            n
491        } else {
492            (i + 1) * n / num_buckets
493        };
494        (start, end)
495    };
496
497    let cutoffs: Vec<f32> = (1..num_buckets)
498        .map(|i| sorted[i * n / num_buckets])
499        .collect();
500
501    let weights: Vec<f32> = (0..num_buckets)
502        .map(|i| {
503            let (start, end) = bucket_bounds(i);
504            // If the bucket is empty (e.g., many duplicate values pushed
505            // everyone into one slice), fall back to the nearest real
506            // sample so the decoder still has a sensible value.
507            if start == end {
508                let idx = start.min(n - 1);
509                sorted[idx]
510            } else {
511                let slice = &sorted[start..end];
512                slice.iter().sum::<f32>() / slice.len() as f32
513            }
514        })
515        .collect();
516
517    (cutoffs, weights)
518}
519
520/// Return the index of the bucket that `value` falls into given a set of
521/// ascending cutoffs.
522///
523/// Values strictly below the first cutoff go into bucket 0; values at or
524/// above the last cutoff go into the top bucket (`cutoffs.len()`). This
525/// matches the "lower-inclusive" convention used throughout PLAID.
526fn bucket_for_value(value: f32, cutoffs: &[f32]) -> u8 {
527    let mut idx = 0u8;
528    for cutoff in cutoffs {
529        if value >= *cutoff {
530            idx += 1;
531        } else {
532            break;
533        }
534    }
535    idx
536}
537
538#[cfg(test)]
539mod tests {
540    use super::*;
541
542    /// Build a minimal 2-bit codec over 1-D residuals with symmetric
543    /// cutoffs around zero, handy for checking encode/decode without
544    /// worrying about centroid geometry.
545    fn two_bit_1d_codec_with_centroids(centroids: Vec<f32>) -> ResidualCodec {
546        ResidualCodec {
547            nbits: 2,
548            dim: 1,
549            centroids,
550            bucket_cutoffs: vec![-0.5, 0.0, 0.5],
551            bucket_weights: vec![-0.75, -0.25, 0.25, 0.75],
552        }
553    }
554
555    #[test]
556    fn decode_with_lookup_table_matches_scalar_decode() {
557        // Paper §4.5: precompute the 2^8 possible unpack outputs for a
558        // packed byte, decode via table lookup instead of bit ops. The
559        // output must match the scalar reference bit-for-bit.
560        for &nbits in &[1u32, 2, 4, 8] {
561            let num_buckets = 1usize << nbits;
562            let codec = ResidualCodec {
563                nbits,
564                dim: 16,
565                centroids: (0..16).map(|i| i as f32 * 0.1).collect(),
566                bucket_cutoffs: (1..num_buckets)
567                    .map(|i| (i as f32 / num_buckets as f32) - 0.5)
568                    .collect(),
569                bucket_weights: (0..num_buckets)
570                    .map(|i| (i as f32 + 0.5) / num_buckets as f32 - 0.5)
571                    .collect(),
572            };
573            let input: Vec<f32> =
574                (0..16).map(|i| i as f32 * 0.05 - 0.3).collect();
575            let encoded = codec.encode_vector(&input).unwrap();
576
577            let scalar = codec.decode_vector(&encoded).unwrap();
578            let table = DecodeTable::new(&codec);
579            let via_table =
580                codec.decode_vector_with_table(&encoded, &table).unwrap();
581            assert_eq!(scalar, via_table, "mismatch at nbits={nbits}");
582        }
583    }
584
585    #[test]
586    fn pack_then_read_code_recovers_every_input() {
587        // Every nbits ∈ {1,2,4,8} should pack losslessly: reading each
588        // position back from the packed buffer must return the
589        // original value.
590        for &nbits in &[1u32, 2, 4, 8] {
591            let num_buckets = 1usize << nbits;
592            // Cycle through 0..num_buckets so every code value lands
593            // somewhere, plus a few more for good byte alignment.
594            let unpacked: Vec<u8> = (0..32u8)
595                .map(|i| (i as usize % num_buckets) as u8)
596                .collect();
597            let packed = pack_codes(&unpacked, nbits);
598            for (i, &expected) in unpacked.iter().enumerate() {
599                let got = read_code(&packed, i, nbits);
600                assert_eq!(
601                    got, expected,
602                    "nbits={nbits} position {i}: got {got}, expected {expected}",
603                );
604            }
605            // Byte count matches the advertised formula.
606            assert_eq!(
607                packed.len(),
608                packed_bytes_per_vector(unpacked.len(), nbits),
609            );
610        }
611    }
612
613    #[test]
614    fn encode_vector_produces_packed_codes_at_two_bits() {
615        // Paper §4.5: ColBERTv2/PLAID pack `8/nbits` residual codes
616        // per byte. For dim=8 at 2-bit, that's 4 codes per byte ⇒
617        // 2 bytes of packed storage, not 8.
618        let codec = ResidualCodec {
619            nbits: 2,
620            dim: 8,
621            centroids: vec![0.0; 8],
622            bucket_cutoffs: vec![-0.5, 0.0, 0.5],
623            bucket_weights: vec![-0.75, -0.25, 0.25, 0.75],
624        };
625        let encoded = codec.encode_vector(&[0.1f32; 8]).unwrap();
626        assert_eq!(encoded.codes.len(), 2);
627    }
628
629    #[test]
630    fn encode_vector_produces_packed_codes_at_four_bits() {
631        // dim=8 at 4-bit ⇒ 2 codes per byte ⇒ 4 bytes.
632        let codec = ResidualCodec {
633            nbits: 4,
634            dim: 8,
635            centroids: vec![0.0; 8],
636            bucket_cutoffs: (0..15).map(|i| i as f32 / 15.0 - 0.5).collect(),
637            bucket_weights: (0..16).map(|i| i as f32 / 16.0 - 0.5).collect(),
638        };
639        let encoded = codec.encode_vector(&[0.1f32; 8]).unwrap();
640        assert_eq!(encoded.codes.len(), 4);
641    }
642
643    #[test]
644    fn encode_decode_roundtrip_at_every_supported_nbits() {
645        // Roundtripping through packing/unpacking must be lossless up
646        // to the bucket quantisation. We exercise 1, 2, 4, and 8 bit
647        // widths on a small residual so every branch of the packing
648        // math gets hit.
649        for nbits in [1u32, 2, 4, 8] {
650            let num_buckets = 1usize << nbits;
651            let bucket_cutoffs: Vec<f32> = (1..num_buckets)
652                .map(|i| (i as f32 / num_buckets as f32) - 0.5)
653                .collect();
654            let bucket_weights: Vec<f32> = (0..num_buckets)
655                .map(|i| (i as f32 + 0.5) / num_buckets as f32 - 0.5)
656                .collect();
657            let codec = ResidualCodec {
658                nbits,
659                dim: 8,
660                centroids: vec![0.0; 8],
661                bucket_cutoffs,
662                bucket_weights,
663            };
664            let input = [-0.4f32, -0.1, 0.0, 0.25, 0.49, -0.25, 0.1, 0.3];
665            let encoded = codec.encode_vector(&input).unwrap();
666            let decoded = codec.decode_vector(&encoded).unwrap();
667            let max_err = input
668                .iter()
669                .zip(decoded.iter())
670                .map(|(a, b)| (a - b).abs())
671                .fold(0.0f32, f32::max);
672            // Each bucket covers at most `1 / num_buckets` of [−0.5, 0.5]
673            // so reconstruction error per dim is bounded by half a
674            // bucket width.
675            let tolerance = 1.0 / num_buckets as f32;
676            assert!(
677                max_err <= tolerance,
678                "nbits={nbits}: max_err={max_err}, tolerance={tolerance}",
679            );
680        }
681    }
682
683    #[test]
684    fn bucket_for_value_places_below_first_cutoff_in_bucket_zero() {
685        let cutoffs = [-0.5, 0.0, 0.5];
686        assert_eq!(bucket_for_value(-1.0, &cutoffs), 0);
687    }
688
689    #[test]
690    fn bucket_for_value_places_at_or_above_last_cutoff_in_top_bucket() {
691        let cutoffs = [-0.5, 0.0, 0.5];
692        assert_eq!(bucket_for_value(0.5, &cutoffs), 3);
693        assert_eq!(bucket_for_value(9.9, &cutoffs), 3);
694    }
695
696    #[test]
697    fn bucket_for_value_picks_intermediate_buckets() {
698        let cutoffs = [-0.5, 0.0, 0.5];
699        assert_eq!(bucket_for_value(-0.25, &cutoffs), 1);
700        assert_eq!(bucket_for_value(0.25, &cutoffs), 2);
701    }
702
703    #[test]
704    fn num_buckets_is_two_to_the_nbits() {
705        let codec = two_bit_1d_codec_with_centroids(vec![0.0]);
706        assert_eq!(codec.num_buckets(), 4);
707
708        let mut four_bit = codec.clone();
709        four_bit.nbits = 4;
710        four_bit.bucket_cutoffs = (0..15).map(|i| i as f32 / 15.0).collect();
711        four_bit.bucket_weights = (0..16).map(|i| i as f32).collect();
712        assert_eq!(four_bit.num_buckets(), 16);
713    }
714
715    #[test]
716    fn encode_picks_nearest_centroid() {
717        // Two 1-D centroids at 0 and 10. Input 9.0 should snap to
718        // centroid 1 (distance 1) rather than centroid 0 (distance 9).
719        let codec = two_bit_1d_codec_with_centroids(vec![0.0, 10.0]);
720        let encoded = codec.encode_vector(&[9.0]).unwrap();
721        assert_eq!(encoded.centroid_id, 1);
722    }
723
724    #[test]
725    fn decode_inverts_a_known_encoding() {
726        let codec = two_bit_1d_codec_with_centroids(vec![0.0]);
727        // Residual −0.3 → bucket 1 (between −0.5 and 0.0) → weight −0.25.
728        let encoded = codec.encode_vector(&[-0.3]).unwrap();
729        assert_eq!(encoded.codes, vec![1]);
730        let decoded = codec.decode_vector(&encoded).unwrap();
731        assert_eq!(decoded, vec![-0.25]);
732    }
733
734    #[test]
735    fn encode_then_decode_stays_inside_bucket_half_width() {
736        // With cutoffs [-0.5, 0, 0.5] and weights at bucket midpoints,
737        // reconstruction error per dim is at most 0.25 for any value in
738        // the middle buckets.
739        let codec = two_bit_1d_codec_with_centroids(vec![0.0]);
740        for &value in &[-0.4f32, -0.1, 0.0, 0.2, 0.4] {
741            let encoded = codec.encode_vector(&[value]).unwrap();
742            let decoded = codec.decode_vector(&encoded).unwrap();
743            assert!(
744                (decoded[0] - value).abs() <= 0.25,
745                "value {value} -> decoded {d}",
746                d = decoded[0],
747            );
748        }
749    }
750
751    #[test]
752    fn reconstruction_error_is_zero_when_residual_exactly_matches_weight() {
753        let codec = two_bit_1d_codec_with_centroids(vec![0.0]);
754        // Residual −0.25 lives in bucket 1, which decodes to −0.25.
755        assert_eq!(codec.reconstruction_error(&[-0.25]).unwrap(), 0.0);
756    }
757
758    #[test]
759    fn validate_rejects_wrong_number_of_cutoffs() {
760        let mut codec = two_bit_1d_codec_with_centroids(vec![0.0]);
761        codec.bucket_cutoffs.push(1.0); // now 4 cutoffs, expected 3
762        assert!(codec.validate().is_err());
763    }
764
765    #[test]
766    fn validate_rejects_non_monotonic_cutoffs() {
767        let mut codec = two_bit_1d_codec_with_centroids(vec![0.0]);
768        codec.bucket_cutoffs = vec![0.5, 0.0, 0.5];
769        assert!(codec.validate().is_err());
770    }
771
772    #[test]
773    #[should_panic(expected = "packed bytes")]
774    fn decode_panics_on_wrong_packed_code_length() {
775        // With `dim=1` at 2 bits, a valid encoding is 1 packed byte.
776        // Handing decode a 2-byte buffer should fail the shape check.
777        let codec = two_bit_1d_codec_with_centroids(vec![0.0]);
778        let bad = EncodedVector {
779            centroid_id: 0,
780            codes: vec![0, 0],
781        };
782        let _ = codec.decode_vector(&bad).unwrap();
783    }
784
785    #[test]
786    fn train_quantizer_produces_right_number_of_cutoffs_and_weights() {
787        let residuals: Vec<f32> =
788            (0..1000).map(|i| i as f32 / 1000.0).collect();
789        let (cutoffs, weights) = train_quantizer(&residuals, 2);
790        assert_eq!(cutoffs.len(), 3);
791        assert_eq!(weights.len(), 4);
792    }
793
794    #[test]
795    fn train_quantizer_cutoffs_are_monotonic() {
796        let residuals: Vec<f32> =
797            (0..2048).map(|i| (i as f32 / 2048.0) - 0.5).collect();
798        let (cutoffs, _) = train_quantizer(&residuals, 4);
799        for pair in cutoffs.windows(2) {
800            assert!(
801                pair[0] <= pair[1],
802                "cutoffs must be non-decreasing: {pair:?}"
803            );
804        }
805    }
806
807    #[test]
808    fn train_quantizer_on_uniform_data_gives_quartile_cutoffs() {
809        // Uniform samples in [0, 1000) with 2 bits → quartile cutoffs at
810        // roughly 250, 500, 750.
811        let residuals: Vec<f32> = (0..1000).map(|i| i as f32).collect();
812        let (cutoffs, _) = train_quantizer(&residuals, 2);
813        assert!((cutoffs[0] - 250.0).abs() < 1.0);
814        assert!((cutoffs[1] - 500.0).abs() < 1.0);
815        assert!((cutoffs[2] - 750.0).abs() < 1.0);
816    }
817
818    #[test]
819    fn train_quantizer_weights_bracket_cutoffs() {
820        // Each weight should fall within its bucket's [low, high] range.
821        // For uniform data, this is straightforward to verify.
822        let residuals: Vec<f32> = (0..1024).map(|i| i as f32).collect();
823        let (cutoffs, weights) = train_quantizer(&residuals, 2);
824
825        // Bucket 0: below cutoffs[0]
826        assert!(weights[0] < cutoffs[0]);
827        // Bucket 3: above cutoffs[2]
828        assert!(weights[3] > cutoffs[2]);
829        // Middle buckets fall inside their cutoff ranges.
830        assert!(cutoffs[0] <= weights[1] && weights[1] < cutoffs[1]);
831        assert!(cutoffs[1] <= weights[2] && weights[2] < cutoffs[2]);
832    }
833
834    #[test]
835    #[should_panic(expected = "empty sample")]
836    fn train_quantizer_panics_on_empty_sample() {
837        let _ = train_quantizer(&[], 2);
838    }
839
840    #[test]
841    #[should_panic(expected = "NaN")]
842    fn train_quantizer_panics_on_nan() {
843        let _ = train_quantizer(&[0.1, f32::NAN, 0.3], 2);
844    }
845
846    #[test]
847    fn trained_codec_round_trips_within_reasonable_error() {
848        // Train a 4-bit codec on synthetic residuals, then check the
849        // reconstruction error on held-out samples is small relative to
850        // the residual magnitude.
851        let training: Vec<f32> =
852            (0..2048).map(|i| (i as f32 / 2048.0) - 0.5).collect();
853        let (cutoffs, weights) = train_quantizer(&training, 4);
854
855        let codec = ResidualCodec {
856            nbits: 4,
857            dim: 1,
858            centroids: vec![0.0],
859            bucket_cutoffs: cutoffs,
860            bucket_weights: weights,
861        };
862        codec.validate().unwrap();
863
864        let mut max_err: f32 = 0.0;
865        for v in &[-0.4f32, -0.1, 0.0, 0.25, 0.49] {
866            let err = codec.reconstruction_error(&[*v]).unwrap().sqrt();
867            max_err = max_err.max(err);
868        }
869        // 16 buckets spanning ~1.0 of range ⇒ each bucket ≈ 0.0625 wide,
870        // so reconstruction error should sit well below 0.05.
871        assert!(
872            max_err < 0.05,
873            "max reconstruction error {max_err} above tolerance"
874        );
875    }
876
877    #[test]
878    fn encode_and_decode_roundtrip_multi_dim_stays_close() {
879        // 2-D centroid at (1, 1). For an input (1.1, 0.7), the residual
880        // is (0.1, -0.3). Both land in inner buckets and decode back to
881        // values within 0.25 of the truth per dimension.
882        let codec = ResidualCodec {
883            nbits: 2,
884            dim: 2,
885            centroids: vec![1.0, 1.0],
886            bucket_cutoffs: vec![-0.5, 0.0, 0.5],
887            bucket_weights: vec![-0.75, -0.25, 0.25, 0.75],
888        };
889        let input = [1.1f32, 0.7];
890        let encoded = codec.encode_vector(&input).unwrap();
891        let decoded = codec.decode_vector(&encoded).unwrap();
892        for (d, i) in decoded.iter().zip(input.iter()) {
893            assert!((d - i).abs() <= 0.25);
894        }
895    }
896}