Skip to main content

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}