1use crate::{
30 PlaidError,
31 Result,
32 distance::squared_l2,
33 kmeans::nearest_centroid,
34};
35
36#[derive(Debug, Clone)]
44pub struct ResidualCodec {
45 pub nbits: u32,
47 pub dim: usize,
49 pub centroids: Vec<f32>,
51 pub bucket_cutoffs: Vec<f32>,
53 pub bucket_weights: Vec<f32>,
55}
56
57#[derive(Debug, Clone, PartialEq, Eq)]
67pub struct EncodedVector {
68 pub centroid_id: u32,
70 pub codes: Vec<u8>,
73}
74
75pub 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
90fn 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
109pub 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
122pub struct DecodeTable {
132 weights: Vec<f32>,
135 codes_per_byte: usize,
136 nbits: u32,
137}
138
139impl DecodeTable {
140 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 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 pub fn codes_per_byte(&self) -> usize {
172 self.codes_per_byte
173 }
174
175 pub fn nbits(&self) -> u32 {
177 self.nbits
178 }
179}
180
181impl ResidualCodec {
182 pub fn num_buckets(&self) -> usize {
184 1usize << self.nbits
185 }
186
187 pub fn num_centroids(&self) -> usize {
189 self.centroids.len() / self.dim
190 }
191
192 pub fn packed_bytes(&self) -> usize {
194 packed_bytes_per_vector(self.dim, self.nbits)
195 }
196
197 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 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 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 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 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 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
449pub 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 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 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
520fn 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 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 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 for &nbits in &[1u32, 2, 4, 8] {
591 let num_buckets = 1usize << nbits;
592 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 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 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 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 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 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 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 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 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 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); 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 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 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 let residuals: Vec<f32> = (0..1024).map(|i| i as f32).collect();
823 let (cutoffs, weights) = train_quantizer(&residuals, 2);
824
825 assert!(weights[0] < cutoffs[0]);
827 assert!(weights[3] > cutoffs[2]);
829 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 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 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 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}