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 ¢roids_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 ¢roids_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}