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 candle_core::Tensor;
30
31use crate::{
32    PlaidError,
33    Result,
34    device::default_device,
35    distance::squared_l2,
36    kmeans::{assign_as_tensor, nearest_centroid},
37};
38
39/// Memory budget for the per-chunk residual + bucket tensors during
40/// GPU-batched encoding.
41///
42/// Mirrors [`crate::kmeans::ASSIGN_CHUNK_BYTES`]. A chunk produces a
43/// `[chunk, dim] f32` retrieved-centroids tensor plus a `[chunk, dim]
44/// f32` residuals tensor plus a `[chunk, dim] u32` buckets tensor;
45/// at 128 MiB per block that keeps the working set comfortably under
46/// 512 MiB on top of the resident tokens tensor, leaving headroom for
47/// cuBLAS workspace and the caller's encoder model.
48const ENCODE_CHUNK_BYTES: usize = 128 * 1024 * 1024;
49
50/// Pick a chunk size for batch encoding so the heaviest per-chunk
51/// tensor stays under [`ENCODE_CHUNK_BYTES`].
52///
53/// `dim` determines the f32 residual chunk (`chunk * dim * 4`);
54/// `packed_bytes` is only included to prevent the pack output from
55/// blowing past the budget when nbits is large.
56fn encode_chunk_rows(dim: usize, _packed_bytes: usize) -> usize {
57    let bytes_per_row = dim * std::mem::size_of::<f32>();
58    (ENCODE_CHUNK_BYTES / bytes_per_row).max(1)
59}
60
61/// A trained residual-quantization codec.
62///
63/// Cutoffs partition the real line into `2^nbits` buckets. `bucket_cutoffs`
64/// holds the `2^nbits - 1` internal boundaries in ascending order;
65/// `bucket_weights` holds the `2^nbits` reconstruction values used when
66/// decoding. Both are codec-wide: the same cutoffs/weights are applied to
67/// every residual dimension of every token.
68#[derive(Debug, Clone)]
69pub struct ResidualCodec {
70    /// Number of bits per residual dimension. Typically 2 or 4.
71    pub nbits: u32,
72    /// Dimensionality of the (original) token embeddings.
73    pub dim: usize,
74    /// Flat row-major coarse centroids, `k × dim`.
75    pub centroids: Vec<f32>,
76    /// `(2^nbits) - 1` ascending cutoff values for bucketing residuals.
77    pub bucket_cutoffs: Vec<f32>,
78    /// `2^nbits` reconstruction values, one per bucket.
79    pub bucket_weights: Vec<f32>,
80}
81
82/// A single encoded token: a centroid reference plus a bit-packed
83/// buffer of per-dim bucket codes.
84///
85/// The `codes` buffer holds `dim` quantization codes packed LSB-first
86/// at `nbits` bits each. For the supported widths of 1, 2, 4, and 8
87/// bits the buffer length is `(dim * nbits) / 8` (dim is expected to
88/// be a multiple of `8/nbits` so code positions don't span bytes —
89/// ColBERT dims are 128 or 96, which satisfies that constraint for
90/// every supported `nbits`).
91#[derive(Debug, Clone, PartialEq, Eq)]
92pub struct EncodedVector {
93    /// Index of the coarse centroid this token was quantized against.
94    pub centroid_id: u32,
95    /// Bit-packed bucket codes. Use [`ResidualCodec::read_code`] or the
96    /// codec's `decode_vector` to pull values out.
97    pub codes: Vec<u8>,
98}
99
100/// Number of bytes required to pack `dim` codes at `nbits` bits each.
101///
102/// Panics if `nbits` is not one of the supported packed widths.
103pub fn packed_bytes_per_vector(dim: usize, nbits: u32) -> usize {
104    assert_supported_nbits(nbits);
105    (dim * nbits as usize).div_ceil(8)
106}
107
108fn assert_supported_nbits(nbits: u32) {
109    assert!(
110        matches!(nbits, 1 | 2 | 4 | 8),
111        "packed codec: nbits must be 1, 2, 4, or 8 (got {nbits})",
112    );
113}
114
115/// Pack `unpacked` (one byte per code, values in `0 .. 2^nbits`) into
116/// an LSB-first bit-packed buffer.
117fn pack_codes(unpacked: &[u8], nbits: u32) -> Vec<u8> {
118    assert_supported_nbits(nbits);
119    if nbits == 8 {
120        return unpacked.to_vec();
121    }
122    let codes_per_byte = 8 / nbits as usize;
123    let mask: u8 = ((1u16 << nbits) - 1) as u8;
124    let n_bytes = unpacked.len().div_ceil(codes_per_byte);
125    let mut packed = vec![0u8; n_bytes];
126    for (i, &code) in unpacked.iter().enumerate() {
127        let byte_idx = i / codes_per_byte;
128        let bit_off = (i % codes_per_byte) * nbits as usize;
129        packed[byte_idx] |= (code & mask) << bit_off;
130    }
131    packed
132}
133
134/// Read the code at logical position `i` from a packed buffer.
135pub fn read_code(packed: &[u8], i: usize, nbits: u32) -> u8 {
136    assert_supported_nbits(nbits);
137    if nbits == 8 {
138        return packed[i];
139    }
140    let codes_per_byte = 8 / nbits as usize;
141    let mask: u8 = ((1u16 << nbits) - 1) as u8;
142    let byte_idx = i / codes_per_byte;
143    let bit_off = (i % codes_per_byte) * nbits as usize;
144    (packed[byte_idx] >> bit_off) & mask
145}
146
147/// Precomputed 256-entry lookup table mapping every possible packed
148/// byte to the sequence of `bucket_weights` values it decodes to.
149///
150/// PLAID §4.5 notes that naive decompression pays a chain of
151/// shift-mask-weight-lookup operations per residual dimension; a
152/// one-off table that already composes the shift/mask with the weight
153/// lookup reduces decoding to a single load per code position. For
154/// `nbits=2` the whole table is `256 × 4` f32 = 4 KiB and easily stays
155/// in L1.
156pub struct DecodeTable {
157    /// `weights[b * codes_per_byte + k]` = weight for the `k`-th code
158    /// position inside packed byte value `b`.
159    weights: Vec<f32>,
160    codes_per_byte: usize,
161    nbits: u32,
162}
163
164impl DecodeTable {
165    /// Build the table for `codec`. Call once per search/decode batch
166    /// and reuse across every encoded vector.
167    pub fn new(codec: &ResidualCodec) -> Self {
168        assert_supported_nbits(codec.nbits);
169        let codes_per_byte = 8 / codec.nbits as usize;
170        let entries = 256;
171        let mut weights = vec![0.0f32; entries * codes_per_byte];
172        let mask: u8 = ((1u16 << codec.nbits) - 1) as u8;
173        for b in 0u16..256 {
174            let byte = b as u8;
175            for k in 0..codes_per_byte {
176                let code = (byte >> (k * codec.nbits as usize)) & mask;
177                weights[b as usize * codes_per_byte + k] =
178                    codec.bucket_weights[code as usize];
179            }
180        }
181        Self {
182            weights,
183            codes_per_byte,
184            nbits: codec.nbits,
185        }
186    }
187
188    /// Weights for the `codes_per_byte` positions inside packed byte
189    /// `byte`. Length always equals `codes_per_byte`.
190    pub fn weights_for(&self, byte: u8) -> &[f32] {
191        let start = byte as usize * self.codes_per_byte;
192        &self.weights[start..start + self.codes_per_byte]
193    }
194
195    /// Raw row-major `[256, codes_per_byte]` weights buffer.
196    ///
197    /// Exposed so the search path can upload the table once per query
198    /// and decode residuals via batched `index_select` on the device,
199    /// matching the GPU decompression kernel described in PLAID §4.5
200    /// (one thread per packed byte).
201    pub fn weights_flat(&self) -> &[f32] {
202        &self.weights
203    }
204
205    /// Number of codes packed into one byte at this table's `nbits`.
206    pub fn codes_per_byte(&self) -> usize {
207        self.codes_per_byte
208    }
209
210    /// Bit-width the table was built for.
211    pub fn nbits(&self) -> u32 {
212        self.nbits
213    }
214}
215
216impl ResidualCodec {
217    /// Number of buckets this codec partitions the residual space into.
218    pub fn num_buckets(&self) -> usize {
219        1usize << self.nbits
220    }
221
222    /// Number of coarse centroids stored.
223    pub fn num_centroids(&self) -> usize {
224        self.centroids.len() / self.dim
225    }
226
227    /// Number of packed bytes each encoded vector uses.
228    pub fn packed_bytes(&self) -> usize {
229        packed_bytes_per_vector(self.dim, self.nbits)
230    }
231
232    /// Validate internal shape invariants. Called automatically by
233    /// encode/decode; exposed so callers loading a codec from disk can
234    /// fail fast.
235    ///
236    /// # Errors
237    ///
238    /// Returns [`PlaidError::InvalidCodec`] with a description of the
239    /// constraint that's violated.
240    pub fn validate(&self) -> Result<()> {
241        if self.dim == 0 {
242            return Err(PlaidError::InvalidCodec(
243                "codec: dim must be positive".into(),
244            ));
245        }
246        if !matches!(self.nbits, 1 | 2 | 4 | 8) {
247            return Err(PlaidError::InvalidCodec(format!(
248                "codec: nbits must be 1, 2, 4, or 8, got {}",
249                self.nbits
250            )));
251        }
252        if !self.centroids.len().is_multiple_of(self.dim)
253            || self.centroids.is_empty()
254        {
255            return Err(PlaidError::InvalidCodec(format!(
256                "codec: centroids length {} is not a positive multiple of dim {}",
257                self.centroids.len(),
258                self.dim,
259            )));
260        }
261        let expected_buckets = self.num_buckets();
262        if self.bucket_weights.len() != expected_buckets {
263            return Err(PlaidError::InvalidCodec(format!(
264                "codec: expected {} bucket_weights, got {}",
265                expected_buckets,
266                self.bucket_weights.len(),
267            )));
268        }
269        if self.bucket_cutoffs.len() != expected_buckets - 1 {
270            return Err(PlaidError::InvalidCodec(format!(
271                "codec: expected {} bucket_cutoffs, got {}",
272                expected_buckets - 1,
273                self.bucket_cutoffs.len(),
274            )));
275        }
276        for pair in self.bucket_cutoffs.windows(2) {
277            if pair[0] > pair[1] || pair[0].is_nan() || pair[1].is_nan() {
278                return Err(PlaidError::InvalidCodec(
279                    "codec: bucket_cutoffs must be non-decreasing and finite"
280                        .into(),
281                ));
282            }
283        }
284        Ok(())
285    }
286
287    /// Encode a single token embedding.
288    ///
289    /// Finds the nearest centroid, computes the residual, and quantizes
290    /// each dimension against `bucket_cutoffs`.
291    ///
292    /// # Errors
293    ///
294    /// Returns [`PlaidError::InvalidCodec`] if this codec fails its
295    /// shape invariants.
296    ///
297    /// # Panics
298    ///
299    /// Panics if `vector.len() != dim`.
300    pub fn encode_vector(&self, vector: &[f32]) -> Result<EncodedVector> {
301        self.validate()?;
302        assert_eq!(
303            vector.len(),
304            self.dim,
305            "encode_vector: expected {} dims, got {}",
306            self.dim,
307            vector.len(),
308        );
309
310        let centroid_id = nearest_centroid(vector, &self.centroids, self.dim);
311        let centroid_slice = &self.centroids
312            [centroid_id * self.dim..(centroid_id + 1) * self.dim];
313
314        let unpacked: Vec<u8> = vector
315            .iter()
316            .zip(centroid_slice.iter())
317            .map(|(v, c)| bucket_for_value(*v - *c, &self.bucket_cutoffs))
318            .collect();
319        let codes = pack_codes(&unpacked, self.nbits);
320
321        Ok(EncodedVector {
322            centroid_id: centroid_id as u32,
323            codes,
324        })
325    }
326
327    /// Encode every token in a flat `n × dim` buffer in one batched
328    /// pass, returning the per-token centroid id and a flat `n × dim`
329    /// code buffer.
330    ///
331    /// The expensive step — the nearest-centroid lookup — runs as a
332    /// single matmul through [`crate::kmeans::assign_points`], which
333    /// uses candle's GEMM (CPU or CUDA depending on build). The
334    /// residual + bucket loop stays scalar because per-element
335    /// `searchsorted` would otherwise require either a 3-D broadcast
336    /// against the cutoffs table or a per-cutoff kernel launch — both
337    /// less efficient than a tight Rust loop over the small cutoffs
338    /// vector. Returning the codes flat avoids `n` `Vec<u8>`
339    /// allocations; callers split into per-token slices as needed.
340    ///
341    /// # Errors
342    ///
343    /// Returns [`PlaidError::InvalidCodec`] if the codec fails its
344    /// shape invariants, or [`PlaidError::Tensor`] if the
345    /// matmul-driven nearest-centroid lookup fails.
346    ///
347    /// # Panics
348    ///
349    /// Panics if `tokens.len() % dim != 0` or if `tokens` is empty.
350    ///
351    /// [`PlaidError::InvalidCodec`]: crate::PlaidError::InvalidCodec
352    /// [`PlaidError::Tensor`]: crate::PlaidError::Tensor
353    pub fn batch_encode_tokens(
354        &self,
355        tokens: &[f32],
356    ) -> Result<(Vec<u32>, Vec<u8>)> {
357        let chunk_rows = encode_chunk_rows(self.dim, self.packed_bytes());
358        self.batch_encode_tokens_with_chunk_rows(tokens, chunk_rows)
359    }
360
361    /// Same as [`batch_encode_tokens`] but processes the input in
362    /// tiles of `chunk_rows` tokens per upload.
363    ///
364    /// This is the path the PLAID builder uses for pools that would
365    /// otherwise exceed VRAM — only `[chunk_rows, dim] f32` ever lives
366    /// on the device at once, so peak VRAM stays bounded in
367    /// `chunk_rows` regardless of corpus size or embedding dimension.
368    /// The output is byte-identical to [`batch_encode_tokens`] for
369    /// any `chunk_rows ≥ 1`, which the
370    /// `prop_batch_encode_chunked_is_partition_invariant` hegel
371    /// property checks across shrunk-counterexample shapes.
372    ///
373    /// # Errors
374    ///
375    /// Returns [`PlaidError::InvalidCodec`] on codec-shape violations,
376    /// or [`PlaidError::Tensor`] if any per-chunk allocation or
377    /// matmul fails.
378    ///
379    /// # Panics
380    ///
381    /// Panics if `tokens.len() % dim != 0` or if `chunk_rows == 0`.
382    ///
383    /// [`PlaidError::InvalidCodec`]: crate::PlaidError::InvalidCodec
384    /// [`PlaidError::Tensor`]: crate::PlaidError::Tensor
385    pub fn batch_encode_tokens_with_chunk_rows(
386        &self,
387        tokens: &[f32],
388        chunk_rows: usize,
389    ) -> Result<(Vec<u32>, Vec<u8>)> {
390        assert!(
391            chunk_rows > 0,
392            "batch_encode_tokens_with_chunk_rows: chunk_rows must be positive"
393        );
394        self.validate()?;
395        assert!(
396            tokens.len().is_multiple_of(self.dim),
397            "batch_encode_tokens_with_chunk_rows: tokens length {} is not a multiple of dim {}",
398            tokens.len(),
399            self.dim,
400        );
401        let n = tokens.len() / self.dim;
402        if n == 0 {
403            return Ok((Vec::new(), Vec::new()));
404        }
405
406        let device = default_device();
407        let k = self.num_centroids();
408        let packed_per_token = self.packed_bytes();
409        let codes_per_byte = 8 / self.nbits as usize;
410
411        // Codec state uploads hoisted above the tile loop — every tile
412        // reuses the same `[k, dim]` centroids, cutoffs, and shift
413        // weights, so re-uploading per tile would dominate the runtime
414        // when chunk_rows is small (e.g. at LateOn scale with dim=1536
415        // and a 128 MiB tile budget, we touch ~310 tiles per pool).
416        let centroids_dev =
417            Tensor::from_slice(&self.centroids, (k, self.dim), device)?;
418        let cutoffs_dev = Tensor::from_slice(
419            &self.bucket_cutoffs,
420            (self.bucket_cutoffs.len(),),
421            device,
422        )?;
423        let shift_weights: Vec<u32> = (0..codes_per_byte)
424            .map(|slot| 1u32 << (slot as u32 * self.nbits))
425            .collect();
426        let shifts_dev =
427            Tensor::from_slice(&shift_weights, (1, 1, codes_per_byte), device)?;
428
429        let mut centroid_ids: Vec<u32> = Vec::with_capacity(n);
430        let mut packed_codes: Vec<u8> =
431            Vec::with_capacity(n * packed_per_token);
432
433        // Walk the host slice in `chunk_rows`-sized tiles. Each tile
434        // uploads only its own `[len, dim] f32` block, passes through
435        // the full GPU pipeline (assign → residual → bucketize →
436        // pack), drains centroid ids and packed bytes back to the
437        // host accumulators, and drops. Device peak is `centroids +
438        // cutoffs + shifts + one live tile` — bounded in `chunk_rows`,
439        // independent of corpus size or dim.
440        let mut start = 0usize;
441        while start < n {
442            let len = chunk_rows.min(n - start);
443            let slice = &tokens[start * self.dim..(start + len) * self.dim];
444            let tile_tensor =
445                Tensor::from_slice(slice, (len, self.dim), device)?;
446            let (tile_cids, tile_codes) = self
447                .encode_chunk_on_tensor_with_state(
448                    &tile_tensor,
449                    len,
450                    &centroids_dev,
451                    &cutoffs_dev,
452                    &shifts_dev,
453                    codes_per_byte,
454                    packed_per_token,
455                )?;
456            centroid_ids.extend(tile_cids);
457            packed_codes.extend(tile_codes);
458            start += len;
459        }
460        Ok((centroid_ids, packed_codes))
461    }
462
463    /// Private per-tile pipeline used by
464    /// [`batch_encode_tokens_with_chunk_rows`] and
465    /// [`batch_encode_tokens_on_tensor`]. Takes pre-uploaded codec
466    /// state (centroids, cutoffs, shift weights) so repeated calls in
467    /// the outer tile loop don't re-upload them.
468    #[allow(clippy::too_many_arguments)]
469    fn encode_chunk_on_tensor_with_state(
470        &self,
471        tile: &Tensor,
472        len: usize,
473        centroids_dev: &Tensor,
474        cutoffs_dev: &Tensor,
475        shifts_dev: &Tensor,
476        codes_per_byte: usize,
477        packed_per_token: usize,
478    ) -> Result<(Vec<u32>, Vec<u8>)> {
479        let device = tile.device();
480
481        // Assign → gather per-token centroids → residuals.
482        let assign_chunk = assign_as_tensor(tile, centroids_dev)?;
483        let retrieved = centroids_dev.index_select(&assign_chunk, 0)?;
484        let residuals = tile.sub(&retrieved)?;
485
486        // Bucketize by accumulating `residuals >= cutoff` across
487        // cutoffs. Equivalent to PyTorch/fast-plaid's
488        // `bucketize(right=True)`.
489        let mut buckets =
490            Tensor::zeros((len, self.dim), candle_core::DType::U32, device)?;
491        for i in 0..self.bucket_cutoffs.len() {
492            let cutoff = cutoffs_dev.narrow(0, i, 1)?;
493            let hit = residuals
494                .broadcast_ge(&cutoff)?
495                .to_dtype(candle_core::DType::U32)?;
496            buckets = buckets.add(&hit)?;
497        }
498
499        // Pack `codes_per_byte` consecutive bucket indices per byte.
500        // Pad to `packed_per_token * codes_per_byte` when dim isn't a
501        // clean multiple — zero-padded slots contribute nothing.
502        let padded_dim = packed_per_token * codes_per_byte;
503        let buckets_padded = if padded_dim == self.dim {
504            buckets
505        } else {
506            let pad_len = padded_dim - self.dim;
507            let pad =
508                Tensor::zeros((len, pad_len), candle_core::DType::U32, device)?;
509            Tensor::cat(&[&buckets, &pad], 1)?
510        };
511        let packed_u32 = buckets_padded
512            .reshape((len, packed_per_token, codes_per_byte))?
513            .broadcast_mul(shifts_dev)?
514            .sum(2)?;
515        let packed_u8 = packed_u32.to_dtype(candle_core::DType::U8)?;
516
517        Ok((
518            assign_chunk.to_vec1::<u32>()?,
519            packed_u8.flatten_all()?.to_vec1::<u8>()?,
520        ))
521    }
522
523    /// Same as [`batch_encode_tokens`] but reuses a pre-uploaded tokens
524    /// tensor for the nearest-centroid matmul.
525    ///
526    /// The PLAID builder runs k-means on this same corpus immediately
527    /// before calling batch encode; threading the device tensor
528    /// through saves the second 3.47 GB host→device copy that would
529    /// otherwise collide with the first one in cudarc's caching
530    /// allocator and OOM a 12 GB card.
531    ///
532    /// `tokens` must be the host-side backing buffer for
533    /// `tokens_tensor` — the residual + bit-pack loop is scalar and
534    /// still reads token bytes from the host slice.
535    ///
536    /// # Errors
537    ///
538    /// Returns [`PlaidError::InvalidCodec`] on codec-shape violations,
539    /// or [`PlaidError::Tensor`] if the nearest-centroid matmul fails.
540    ///
541    /// # Panics
542    ///
543    /// Panics if `tokens.len() % dim != 0`.
544    ///
545    /// [`PlaidError::InvalidCodec`]: crate::PlaidError::InvalidCodec
546    /// [`PlaidError::Tensor`]: crate::PlaidError::Tensor
547    pub fn batch_encode_tokens_on_tensor(
548        &self,
549        tokens_tensor: &Tensor,
550        tokens: &[f32],
551    ) -> Result<(Vec<u32>, Vec<u8>)> {
552        self.validate()?;
553        assert!(
554            tokens.len().is_multiple_of(self.dim),
555            "batch_encode_tokens_on_tensor: tokens length {} is not a multiple of dim {}",
556            tokens.len(),
557            self.dim,
558        );
559        let n = tokens.len() / self.dim;
560        if n == 0 {
561            return Ok((Vec::new(), Vec::new()));
562        }
563
564        let device = tokens_tensor.device();
565        let k = self.num_centroids();
566        let packed_per_token = self.packed_bytes();
567        let codes_per_byte = 8 / self.nbits as usize;
568
569        // Codec state uploads hoisted above the tile loop, same as the
570        // host-side [`batch_encode_tokens_with_chunk_rows`] path.
571        let centroids_dev =
572            Tensor::from_slice(&self.centroids, (k, self.dim), device)?;
573        let cutoffs_dev = Tensor::from_slice(
574            &self.bucket_cutoffs,
575            (self.bucket_cutoffs.len(),),
576            device,
577        )?;
578        let shift_weights: Vec<u32> = (0..codes_per_byte)
579            .map(|slot| 1u32 << (slot as u32 * self.nbits))
580            .collect();
581        let shifts_dev =
582            Tensor::from_slice(&shift_weights, (1, 1, codes_per_byte), device)?;
583
584        let mut centroid_ids: Vec<u32> = Vec::with_capacity(n);
585        let mut packed_codes: Vec<u8> =
586            Vec::with_capacity(n * packed_per_token);
587
588        // Chunk over `narrow` views of the pre-uploaded tensor so the
589        // transient residual / buckets / packed tensors stay within
590        // the chunk budget regardless of the full tensor's size.
591        let chunk_rows =
592            encode_chunk_rows(self.dim, packed_per_token).min(n).max(1);
593        let mut start = 0usize;
594        while start < n {
595            let len = chunk_rows.min(n - start);
596            let tile = tokens_tensor.narrow(0, start, len)?;
597            let (tile_cids, tile_codes) = self
598                .encode_chunk_on_tensor_with_state(
599                    &tile,
600                    len,
601                    &centroids_dev,
602                    &cutoffs_dev,
603                    &shifts_dev,
604                    codes_per_byte,
605                    packed_per_token,
606                )?;
607            centroid_ids.extend(tile_cids);
608            packed_codes.extend(tile_codes);
609            start += len;
610        }
611
612        Ok((centroid_ids, packed_codes))
613    }
614
615    /// Reconstruct an approximate token embedding from its codes.
616    ///
617    /// # Errors
618    ///
619    /// Returns [`PlaidError::InvalidCodec`] if this codec fails its
620    /// shape invariants.
621    ///
622    /// # Panics
623    ///
624    /// Panics if `codes.len() != dim` or if any code is out of range.
625    ///
626    /// [`PlaidError::InvalidCodec`]: crate::PlaidError::InvalidCodec
627    pub fn decode_vector(&self, encoded: &EncodedVector) -> Result<Vec<f32>> {
628        let table = DecodeTable::new(self);
629        self.decode_vector_with_table(encoded, &table)
630    }
631
632    /// Decode an encoded vector using a pre-built [`DecodeTable`].
633    ///
634    /// Callers that decode many vectors in a row (e.g., the search
635    /// path's per-candidate decode loop) should build the table once
636    /// outside the loop and reuse it here — each call then amounts to
637    /// one table load per packed byte plus the centroid add.
638    ///
639    /// # Errors
640    ///
641    /// Returns [`PlaidError::InvalidCodec`] if this codec fails its
642    /// shape invariants.
643    ///
644    /// [`PlaidError::InvalidCodec`]: crate::PlaidError::InvalidCodec
645    pub fn decode_vector_with_table(
646        &self,
647        encoded: &EncodedVector,
648        table: &DecodeTable,
649    ) -> Result<Vec<f32>> {
650        self.validate()?;
651        assert_eq!(
652            table.nbits, self.nbits,
653            "decode_vector_with_table: table nbits {} != codec nbits {}",
654            table.nbits, self.nbits,
655        );
656        let expected_bytes = self.packed_bytes();
657        assert_eq!(
658            encoded.codes.len(),
659            expected_bytes,
660            "decode_vector_with_table: expected {expected_bytes} packed bytes, got {}",
661            encoded.codes.len(),
662        );
663        let centroid_id = encoded.centroid_id as usize;
664        assert!(
665            centroid_id < self.num_centroids(),
666            "decode_vector_with_table: centroid_id {} out of range 0..{}",
667            centroid_id,
668            self.num_centroids(),
669        );
670
671        let centroid_slice = &self.centroids
672            [centroid_id * self.dim..(centroid_id + 1) * self.dim];
673        let codes_per_byte = table.codes_per_byte;
674
675        let mut out = Vec::with_capacity(self.dim);
676        for (byte_idx, &byte) in encoded.codes.iter().enumerate() {
677            let weights = table.weights_for(byte);
678            let base_dim = byte_idx * codes_per_byte;
679            for (k, &w) in weights.iter().enumerate() {
680                let dim_idx = base_dim + k;
681                if dim_idx >= self.dim {
682                    break;
683                }
684                out.push(centroid_slice[dim_idx] + w);
685            }
686        }
687        Ok(out)
688    }
689
690    /// Return the squared L2 reconstruction error for `vector` under
691    /// this codec. Useful as a lightweight codec-quality probe in tests
692    /// and evaluation scripts.
693    ///
694    /// # Errors
695    ///
696    /// Returns [`PlaidError::InvalidCodec`] if this codec fails its
697    /// shape invariants.
698    ///
699    /// [`PlaidError::InvalidCodec`]: crate::PlaidError::InvalidCodec
700    pub fn reconstruction_error(&self, vector: &[f32]) -> Result<f32> {
701        let encoded = self.encode_vector(vector)?;
702        let decoded = self.decode_vector(&encoded)?;
703        Ok(squared_l2(vector, &decoded))
704    }
705}
706
707/// Learn bucket cutoffs and reconstruction weights from a sample of
708/// residual values.
709///
710/// The returned tuple is `(bucket_cutoffs, bucket_weights)` with
711/// `2^nbits - 1` cutoffs and `2^nbits` weights, ready to plug into a
712/// [`ResidualCodec`]. Buckets are equal-quantile slices of the input:
713/// cutoffs are picked at `i / (2^nbits)` quantile positions, and each
714/// weight is the arithmetic mean of the residuals falling into that
715/// bucket. This matches fast-plaid's "fit" step and keeps the codec
716/// unbiased on the training distribution.
717///
718/// `residuals` is consumed and sorted in place, so the caller avoids
719/// the ~n·4-byte allocation a borrowed-slice version would need for an
720/// internal copy. NaN values are rejected up front since they would
721/// poison the sort order.
722///
723/// # Panics
724///
725/// Panics if `residuals` is empty, if `nbits` is zero or exceeds 8, or
726/// if `residuals` contains NaN.
727pub fn train_quantizer(
728    mut residuals: Vec<f32>,
729    nbits: u32,
730) -> (Vec<f32>, Vec<f32>) {
731    assert!(!residuals.is_empty(), "train_quantizer: empty sample");
732    assert!(
733        nbits > 0 && nbits <= 8,
734        "train_quantizer: nbits must be in 1..=8, got {nbits}"
735    );
736    assert!(
737        residuals.iter().all(|v| !v.is_nan()),
738        "train_quantizer: residual sample contains NaN"
739    );
740
741    let num_buckets = 1usize << nbits;
742    let n = residuals.len();
743
744    // NaN was rejected above, so `total_cmp` is a strict ordering.
745    // Sort in place to avoid a duplicate copy of the residuals buffer —
746    // on a large corpus this single copy was worth several GB of RSS.
747    residuals.sort_unstable_by(|a, b| a.total_cmp(b));
748
749    let bucket_bounds = |i: usize| -> (usize, usize) {
750        let start = i * n / num_buckets;
751        let end = if i + 1 == num_buckets {
752            n
753        } else {
754            (i + 1) * n / num_buckets
755        };
756        (start, end)
757    };
758
759    let cutoffs: Vec<f32> = (1..num_buckets)
760        .map(|i| residuals[i * n / num_buckets])
761        .collect();
762
763    let weights: Vec<f32> = (0..num_buckets)
764        .map(|i| {
765            let (start, end) = bucket_bounds(i);
766            // If the bucket is empty (e.g., many duplicate values pushed
767            // everyone into one slice), fall back to the nearest real
768            // sample so the decoder still has a sensible value.
769            if start == end {
770                let idx = start.min(n - 1);
771                residuals[idx]
772            } else {
773                let slice = &residuals[start..end];
774                slice.iter().sum::<f32>() / slice.len() as f32
775            }
776        })
777        .collect();
778
779    (cutoffs, weights)
780}
781
782/// Return the index of the bucket that `value` falls into given a set of
783/// ascending cutoffs.
784///
785/// Values strictly below the first cutoff go into bucket 0; values at or
786/// above the last cutoff go into the top bucket (`cutoffs.len()`). This
787/// matches the "lower-inclusive" convention used throughout PLAID.
788fn bucket_for_value(value: f32, cutoffs: &[f32]) -> u8 {
789    let mut idx = 0u8;
790    for cutoff in cutoffs {
791        if value >= *cutoff {
792            idx += 1;
793        } else {
794            break;
795        }
796    }
797    idx
798}
799
800#[cfg(test)]
801mod tests {
802    use super::*;
803
804    /// Build a minimal 2-bit codec over 1-D residuals with symmetric
805    /// cutoffs around zero, handy for checking encode/decode without
806    /// worrying about centroid geometry.
807    fn two_bit_1d_codec_with_centroids(centroids: Vec<f32>) -> ResidualCodec {
808        ResidualCodec {
809            nbits: 2,
810            dim: 1,
811            centroids,
812            bucket_cutoffs: vec![-0.5, 0.0, 0.5],
813            bucket_weights: vec![-0.75, -0.25, 0.25, 0.75],
814        }
815    }
816
817    #[test]
818    fn decode_with_lookup_table_matches_scalar_decode() {
819        // Paper §4.5: precompute the 2^8 possible unpack outputs for a
820        // packed byte, decode via table lookup instead of bit ops. The
821        // output must match the scalar reference bit-for-bit.
822        for &nbits in &[1u32, 2, 4, 8] {
823            let num_buckets = 1usize << nbits;
824            let codec = ResidualCodec {
825                nbits,
826                dim: 16,
827                centroids: (0..16).map(|i| i as f32 * 0.1).collect(),
828                bucket_cutoffs: (1..num_buckets)
829                    .map(|i| (i as f32 / num_buckets as f32) - 0.5)
830                    .collect(),
831                bucket_weights: (0..num_buckets)
832                    .map(|i| (i as f32 + 0.5) / num_buckets as f32 - 0.5)
833                    .collect(),
834            };
835            let input: Vec<f32> =
836                (0..16).map(|i| i as f32 * 0.05 - 0.3).collect();
837            let encoded = codec.encode_vector(&input).unwrap();
838
839            let scalar = codec.decode_vector(&encoded).unwrap();
840            let table = DecodeTable::new(&codec);
841            let via_table =
842                codec.decode_vector_with_table(&encoded, &table).unwrap();
843            assert_eq!(scalar, via_table, "mismatch at nbits={nbits}");
844        }
845    }
846
847    #[test]
848    fn pack_then_read_code_recovers_every_input() {
849        // Every nbits ∈ {1,2,4,8} should pack losslessly: reading each
850        // position back from the packed buffer must return the
851        // original value.
852        for &nbits in &[1u32, 2, 4, 8] {
853            let num_buckets = 1usize << nbits;
854            // Cycle through 0..num_buckets so every code value lands
855            // somewhere, plus a few more for good byte alignment.
856            let unpacked: Vec<u8> = (0..32u8)
857                .map(|i| (i as usize % num_buckets) as u8)
858                .collect();
859            let packed = pack_codes(&unpacked, nbits);
860            for (i, &expected) in unpacked.iter().enumerate() {
861                let got = read_code(&packed, i, nbits);
862                assert_eq!(
863                    got, expected,
864                    "nbits={nbits} position {i}: got {got}, expected {expected}",
865                );
866            }
867            // Byte count matches the advertised formula.
868            assert_eq!(
869                packed.len(),
870                packed_bytes_per_vector(unpacked.len(), nbits),
871            );
872        }
873    }
874
875    #[test]
876    fn encode_vector_produces_packed_codes_at_two_bits() {
877        // Paper §4.5: ColBERTv2/PLAID pack `8/nbits` residual codes
878        // per byte. For dim=8 at 2-bit, that's 4 codes per byte ⇒
879        // 2 bytes of packed storage, not 8.
880        let codec = ResidualCodec {
881            nbits: 2,
882            dim: 8,
883            centroids: vec![0.0; 8],
884            bucket_cutoffs: vec![-0.5, 0.0, 0.5],
885            bucket_weights: vec![-0.75, -0.25, 0.25, 0.75],
886        };
887        let encoded = codec.encode_vector(&[0.1f32; 8]).unwrap();
888        assert_eq!(encoded.codes.len(), 2);
889    }
890
891    #[test]
892    fn encode_vector_produces_packed_codes_at_four_bits() {
893        // dim=8 at 4-bit ⇒ 2 codes per byte ⇒ 4 bytes.
894        let codec = ResidualCodec {
895            nbits: 4,
896            dim: 8,
897            centroids: vec![0.0; 8],
898            bucket_cutoffs: (0..15).map(|i| i as f32 / 15.0 - 0.5).collect(),
899            bucket_weights: (0..16).map(|i| i as f32 / 16.0 - 0.5).collect(),
900        };
901        let encoded = codec.encode_vector(&[0.1f32; 8]).unwrap();
902        assert_eq!(encoded.codes.len(), 4);
903    }
904
905    #[test]
906    fn encode_decode_roundtrip_at_every_supported_nbits() {
907        // Roundtripping through packing/unpacking must be lossless up
908        // to the bucket quantisation. We exercise 1, 2, 4, and 8 bit
909        // widths on a small residual so every branch of the packing
910        // math gets hit.
911        for nbits in [1u32, 2, 4, 8] {
912            let num_buckets = 1usize << nbits;
913            let bucket_cutoffs: Vec<f32> = (1..num_buckets)
914                .map(|i| (i as f32 / num_buckets as f32) - 0.5)
915                .collect();
916            let bucket_weights: Vec<f32> = (0..num_buckets)
917                .map(|i| (i as f32 + 0.5) / num_buckets as f32 - 0.5)
918                .collect();
919            let codec = ResidualCodec {
920                nbits,
921                dim: 8,
922                centroids: vec![0.0; 8],
923                bucket_cutoffs,
924                bucket_weights,
925            };
926            let input = [-0.4f32, -0.1, 0.0, 0.25, 0.49, -0.25, 0.1, 0.3];
927            let encoded = codec.encode_vector(&input).unwrap();
928            let decoded = codec.decode_vector(&encoded).unwrap();
929            let max_err = input
930                .iter()
931                .zip(decoded.iter())
932                .map(|(a, b)| (a - b).abs())
933                .fold(0.0f32, f32::max);
934            // Each bucket covers at most `1 / num_buckets` of [−0.5, 0.5]
935            // so reconstruction error per dim is bounded by half a
936            // bucket width.
937            let tolerance = 1.0 / num_buckets as f32;
938            assert!(
939                max_err <= tolerance,
940                "nbits={nbits}: max_err={max_err}, tolerance={tolerance}",
941            );
942        }
943    }
944
945    #[test]
946    fn bucket_for_value_places_below_first_cutoff_in_bucket_zero() {
947        let cutoffs = [-0.5, 0.0, 0.5];
948        assert_eq!(bucket_for_value(-1.0, &cutoffs), 0);
949    }
950
951    #[test]
952    fn bucket_for_value_places_at_or_above_last_cutoff_in_top_bucket() {
953        let cutoffs = [-0.5, 0.0, 0.5];
954        assert_eq!(bucket_for_value(0.5, &cutoffs), 3);
955        assert_eq!(bucket_for_value(9.9, &cutoffs), 3);
956    }
957
958    #[test]
959    fn bucket_for_value_picks_intermediate_buckets() {
960        let cutoffs = [-0.5, 0.0, 0.5];
961        assert_eq!(bucket_for_value(-0.25, &cutoffs), 1);
962        assert_eq!(bucket_for_value(0.25, &cutoffs), 2);
963    }
964
965    #[test]
966    fn num_buckets_is_two_to_the_nbits() {
967        let codec = two_bit_1d_codec_with_centroids(vec![0.0]);
968        assert_eq!(codec.num_buckets(), 4);
969
970        let mut four_bit = codec.clone();
971        four_bit.nbits = 4;
972        four_bit.bucket_cutoffs = (0..15).map(|i| i as f32 / 15.0).collect();
973        four_bit.bucket_weights = (0..16).map(|i| i as f32).collect();
974        assert_eq!(four_bit.num_buckets(), 16);
975    }
976
977    #[test]
978    fn encode_picks_nearest_centroid() {
979        // Two 1-D centroids at 0 and 10. Input 9.0 should snap to
980        // centroid 1 (distance 1) rather than centroid 0 (distance 9).
981        let codec = two_bit_1d_codec_with_centroids(vec![0.0, 10.0]);
982        let encoded = codec.encode_vector(&[9.0]).unwrap();
983        assert_eq!(encoded.centroid_id, 1);
984    }
985
986    #[test]
987    fn decode_inverts_a_known_encoding() {
988        let codec = two_bit_1d_codec_with_centroids(vec![0.0]);
989        // Residual −0.3 → bucket 1 (between −0.5 and 0.0) → weight −0.25.
990        let encoded = codec.encode_vector(&[-0.3]).unwrap();
991        assert_eq!(encoded.codes, vec![1]);
992        let decoded = codec.decode_vector(&encoded).unwrap();
993        assert_eq!(decoded, vec![-0.25]);
994    }
995
996    #[test]
997    fn encode_then_decode_stays_inside_bucket_half_width() {
998        // With cutoffs [-0.5, 0, 0.5] and weights at bucket midpoints,
999        // reconstruction error per dim is at most 0.25 for any value in
1000        // the middle buckets.
1001        let codec = two_bit_1d_codec_with_centroids(vec![0.0]);
1002        for &value in &[-0.4f32, -0.1, 0.0, 0.2, 0.4] {
1003            let encoded = codec.encode_vector(&[value]).unwrap();
1004            let decoded = codec.decode_vector(&encoded).unwrap();
1005            assert!(
1006                (decoded[0] - value).abs() <= 0.25,
1007                "value {value} -> decoded {d}",
1008                d = decoded[0],
1009            );
1010        }
1011    }
1012
1013    #[test]
1014    fn reconstruction_error_is_zero_when_residual_exactly_matches_weight() {
1015        let codec = two_bit_1d_codec_with_centroids(vec![0.0]);
1016        // Residual −0.25 lives in bucket 1, which decodes to −0.25.
1017        assert_eq!(codec.reconstruction_error(&[-0.25]).unwrap(), 0.0);
1018    }
1019
1020    #[test]
1021    fn validate_rejects_wrong_number_of_cutoffs() {
1022        let mut codec = two_bit_1d_codec_with_centroids(vec![0.0]);
1023        codec.bucket_cutoffs.push(1.0); // now 4 cutoffs, expected 3
1024        assert!(codec.validate().is_err());
1025    }
1026
1027    #[test]
1028    fn validate_rejects_non_monotonic_cutoffs() {
1029        let mut codec = two_bit_1d_codec_with_centroids(vec![0.0]);
1030        codec.bucket_cutoffs = vec![0.5, 0.0, 0.5];
1031        assert!(codec.validate().is_err());
1032    }
1033
1034    #[test]
1035    #[should_panic(expected = "packed bytes")]
1036    fn decode_panics_on_wrong_packed_code_length() {
1037        // With `dim=1` at 2 bits, a valid encoding is 1 packed byte.
1038        // Handing decode a 2-byte buffer should fail the shape check.
1039        let codec = two_bit_1d_codec_with_centroids(vec![0.0]);
1040        let bad = EncodedVector {
1041            centroid_id: 0,
1042            codes: vec![0, 0],
1043        };
1044        let _ = codec.decode_vector(&bad).unwrap();
1045    }
1046
1047    #[test]
1048    fn train_quantizer_produces_right_number_of_cutoffs_and_weights() {
1049        let residuals: Vec<f32> =
1050            (0..1000).map(|i| i as f32 / 1000.0).collect();
1051        let (cutoffs, weights) = train_quantizer(residuals, 2);
1052        assert_eq!(cutoffs.len(), 3);
1053        assert_eq!(weights.len(), 4);
1054    }
1055
1056    #[test]
1057    fn train_quantizer_cutoffs_are_monotonic() {
1058        let residuals: Vec<f32> =
1059            (0..2048).map(|i| (i as f32 / 2048.0) - 0.5).collect();
1060        let (cutoffs, _) = train_quantizer(residuals, 4);
1061        for pair in cutoffs.windows(2) {
1062            assert!(
1063                pair[0] <= pair[1],
1064                "cutoffs must be non-decreasing: {pair:?}"
1065            );
1066        }
1067    }
1068
1069    #[test]
1070    fn train_quantizer_on_uniform_data_gives_quartile_cutoffs() {
1071        // Uniform samples in [0, 1000) with 2 bits → quartile cutoffs at
1072        // roughly 250, 500, 750.
1073        let residuals: Vec<f32> = (0..1000).map(|i| i as f32).collect();
1074        let (cutoffs, _) = train_quantizer(residuals, 2);
1075        assert!((cutoffs[0] - 250.0).abs() < 1.0);
1076        assert!((cutoffs[1] - 500.0).abs() < 1.0);
1077        assert!((cutoffs[2] - 750.0).abs() < 1.0);
1078    }
1079
1080    #[test]
1081    fn train_quantizer_weights_bracket_cutoffs() {
1082        // Each weight should fall within its bucket's [low, high] range.
1083        // For uniform data, this is straightforward to verify.
1084        let residuals: Vec<f32> = (0..1024).map(|i| i as f32).collect();
1085        let (cutoffs, weights) = train_quantizer(residuals, 2);
1086
1087        // Bucket 0: below cutoffs[0]
1088        assert!(weights[0] < cutoffs[0]);
1089        // Bucket 3: above cutoffs[2]
1090        assert!(weights[3] > cutoffs[2]);
1091        // Middle buckets fall inside their cutoff ranges.
1092        assert!(cutoffs[0] <= weights[1] && weights[1] < cutoffs[1]);
1093        assert!(cutoffs[1] <= weights[2] && weights[2] < cutoffs[2]);
1094    }
1095
1096    #[test]
1097    #[should_panic(expected = "empty sample")]
1098    fn train_quantizer_panics_on_empty_sample() {
1099        let _ = train_quantizer(Vec::new(), 2);
1100    }
1101
1102    #[test]
1103    #[should_panic(expected = "NaN")]
1104    fn train_quantizer_panics_on_nan() {
1105        let _ = train_quantizer(vec![0.1, f32::NAN, 0.3], 2);
1106    }
1107
1108    #[test]
1109    fn trained_codec_round_trips_within_reasonable_error() {
1110        // Train a 4-bit codec on synthetic residuals, then check the
1111        // reconstruction error on held-out samples is small relative to
1112        // the residual magnitude.
1113        let training: Vec<f32> =
1114            (0..2048).map(|i| (i as f32 / 2048.0) - 0.5).collect();
1115        let (cutoffs, weights) = train_quantizer(training, 4);
1116
1117        let codec = ResidualCodec {
1118            nbits: 4,
1119            dim: 1,
1120            centroids: vec![0.0],
1121            bucket_cutoffs: cutoffs,
1122            bucket_weights: weights,
1123        };
1124        codec.validate().unwrap();
1125
1126        let mut max_err: f32 = 0.0;
1127        for v in &[-0.4f32, -0.1, 0.0, 0.25, 0.49] {
1128            let err = codec.reconstruction_error(&[*v]).unwrap().sqrt();
1129            max_err = max_err.max(err);
1130        }
1131        // 16 buckets spanning ~1.0 of range ⇒ each bucket ≈ 0.0625 wide,
1132        // so reconstruction error should sit well below 0.05.
1133        assert!(
1134            max_err < 0.05,
1135            "max reconstruction error {max_err} above tolerance"
1136        );
1137    }
1138
1139    #[test]
1140    fn encode_and_decode_roundtrip_multi_dim_stays_close() {
1141        // 2-D centroid at (1, 1). For an input (1.1, 0.7), the residual
1142        // is (0.1, -0.3). Both land in inner buckets and decode back to
1143        // values within 0.25 of the truth per dimension.
1144        let codec = ResidualCodec {
1145            nbits: 2,
1146            dim: 2,
1147            centroids: vec![1.0, 1.0],
1148            bucket_cutoffs: vec![-0.5, 0.0, 0.5],
1149            bucket_weights: vec![-0.75, -0.25, 0.25, 0.75],
1150        };
1151        let input = [1.1f32, 0.7];
1152        let encoded = codec.encode_vector(&input).unwrap();
1153        let decoded = codec.decode_vector(&encoded).unwrap();
1154        for (d, i) in decoded.iter().zip(input.iter()) {
1155            assert!((d - i).abs() <= 0.25);
1156        }
1157    }
1158}