maxsim_lut/packing.rs
1//! How a host packs `nbits`-wide bucket indices into bytes.
2//!
3//! This is the one thing a host has to tell the crate about its codec. The
4//! fused table in [`crate::Lut`] is built by asking, for every byte value and
5//! every key position, which bucket that position holds; nothing else about
6//! the host's storage format leaks in.
7
8use crate::Error;
9
10/// Describes a bit-packing layout of `nbits`-wide bucket indices.
11///
12/// Key `k` of a byte is the `k`-th embedding dimension that byte carries, in
13/// dimension order: a token's packed row `bytes[0..dim·nbits/8]` holds dims
14/// `i·keys_per_byte + k` at `(bytes[i], key k)`.
15pub trait Packing {
16 /// Code width: 1, 2, 4 or 8.
17 fn nbits(&self) -> usize;
18
19 /// Bucket index (`0 .. 2^nbits`) stored at key position `key`
20 /// (`0 .. 8/nbits`) of `byte`.
21 fn bucket_index(&self, byte: u8, key: usize) -> usize;
22
23 /// `8 / nbits`.
24 fn keys_per_byte(&self) -> usize {
25 8 / self.nbits()
26 }
27
28 /// Reference packer, the inverse of [`Packing::bucket_index`], for hosts
29 /// that want to produce rows the same way the tests do. Not optimised.
30 fn pack_row(&self, buckets: &[usize], out: &mut [u8]) {
31 let kpb = self.keys_per_byte();
32 assert_eq!(buckets.len() % kpb, 0, "buckets must fill whole bytes");
33 assert!(out.len() >= buckets.len() / kpb, "output too short");
34 for (i, chunk) in buckets.chunks(kpb).enumerate() {
35 // Search the byte whose expansion matches; 256 candidates, and
36 // this is a reference path, so brute force is fine.
37 let byte = (0..=255u8)
38 .find(|&b| (0..kpb).all(|k| self.bucket_index(b, k) == chunk[k]))
39 .unwrap_or_else(|| {
40 panic!(
41 "no byte encodes buckets {chunk:?}: each must be < 2^nbits, and \
42 Packing::bucket_index must be a bijection over bytes"
43 )
44 });
45 out[i] = byte;
46 }
47 }
48}
49
50/// The ColBERT / PLAID residual layout, shared by ColBERTv2's `ResidualCodec`,
51/// PLAID, fast-plaid, next-plaid and WARP.
52///
53/// `quantize_residuals` writes each dimension's bucket bits MSB-first into the
54/// row, *bit 0 of the bucket first*. So within a byte, key `k` occupies bits
55/// `7 - k·nbits` down to `8 - (k+1)·nbits`, and the bucket index is that
56/// `nbits`-wide segment with its bits reversed. The original decoders express
57/// this as a `byte_reversed_bits_map` followed by a group split; this is the
58/// same function written directly.
59#[derive(Debug, Clone, Copy, PartialEq, Eq)]
60pub struct ColbertPacking {
61 nbits: usize,
62}
63
64impl ColbertPacking {
65 /// `nbits` must be 1, 2, 4 or 8.
66 pub fn new(nbits: usize) -> Result<Self, Error> {
67 match nbits {
68 1 | 2 | 4 | 8 => Ok(Self { nbits }),
69 n => Err(Error::NbitsUnsupported(n)),
70 }
71 }
72}
73
74impl Packing for ColbertPacking {
75 fn nbits(&self) -> usize {
76 self.nbits
77 }
78
79 #[inline]
80 fn bucket_index(&self, byte: u8, key: usize) -> usize {
81 let nbits = self.nbits;
82 debug_assert!(key < 8 / nbits);
83 let shift = 8 - nbits * (key + 1);
84 let segment = (byte as usize >> shift) & ((1 << nbits) - 1);
85 // Reverse the nbits-wide segment: the encoder emitted bucket bit 0
86 // at the highest position of the group.
87 let mut rev = 0usize;
88 for b in 0..nbits {
89 if segment & (1 << b) != 0 {
90 rev |= 1 << (nbits - 1 - b);
91 }
92 }
93 rev
94 }
95}
96
97#[cfg(test)]
98mod tests {
99 use super::*;
100
101 /// Independent re-implementation of next-plaid's `quantize_residuals`
102 /// bit loop, so the packing here is pinned to the encoder, not to itself.
103 fn encoder_pack(buckets: &[usize], nbits: usize) -> Vec<u8> {
104 let mut out = vec![0u8; buckets.len() * nbits / 8];
105 let mut bit_idx = 0usize;
106 for &bucket in buckets {
107 for b in 0..nbits {
108 let bit = ((bucket >> b) & 1) as u8;
109 out[bit_idx / 8] |= bit << (7 - (bit_idx % 8));
110 bit_idx += 1;
111 }
112 }
113 out
114 }
115
116 #[test]
117 fn colbert_packing_matches_encoder_bit_loop() {
118 for nbits in [1usize, 2, 4, 8] {
119 let p = ColbertPacking::new(nbits).unwrap();
120 let kpb = p.keys_per_byte();
121 let n = 1usize << nbits;
122 // Every bucket in every key position, plus a pseudo-random row.
123 let mut buckets: Vec<usize> = Vec::new();
124 for b in 0..n {
125 for _ in 0..kpb {
126 buckets.push(b);
127 }
128 }
129 let mut x = 0x9E3779B9u32;
130 for _ in 0..(64 * kpb) {
131 x ^= x << 13;
132 x ^= x >> 17;
133 x ^= x << 5;
134 buckets.push(x as usize % n);
135 }
136 let bytes = encoder_pack(&buckets, nbits);
137 for (d, &want) in buckets.iter().enumerate() {
138 let got = p.bucket_index(bytes[d / kpb], d % kpb);
139 assert_eq!(got, want, "nbits={nbits} dim={d}");
140 }
141 // And the reference packer round-trips.
142 let mut repacked = vec![0u8; bytes.len()];
143 p.pack_row(&buckets, &mut repacked);
144 assert_eq!(repacked, bytes, "nbits={nbits}");
145 }
146 }
147}