use crate::Error;
pub trait Packing {
fn nbits(&self) -> usize;
fn bucket_index(&self, byte: u8, key: usize) -> usize;
fn keys_per_byte(&self) -> usize {
8 / self.nbits()
}
fn pack_row(&self, buckets: &[usize], out: &mut [u8]) {
let kpb = self.keys_per_byte();
assert_eq!(buckets.len() % kpb, 0, "buckets must fill whole bytes");
assert!(out.len() >= buckets.len() / kpb, "output too short");
for (i, chunk) in buckets.chunks(kpb).enumerate() {
let byte = (0..=255u8)
.find(|&b| (0..kpb).all(|k| self.bucket_index(b, k) == chunk[k]))
.unwrap_or_else(|| {
panic!(
"no byte encodes buckets {chunk:?}: each must be < 2^nbits, and \
Packing::bucket_index must be a bijection over bytes"
)
});
out[i] = byte;
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ColbertPacking {
nbits: usize,
}
impl ColbertPacking {
pub fn new(nbits: usize) -> Result<Self, Error> {
match nbits {
1 | 2 | 4 | 8 => Ok(Self { nbits }),
n => Err(Error::NbitsUnsupported(n)),
}
}
}
impl Packing for ColbertPacking {
fn nbits(&self) -> usize {
self.nbits
}
#[inline]
fn bucket_index(&self, byte: u8, key: usize) -> usize {
let nbits = self.nbits;
debug_assert!(key < 8 / nbits);
let shift = 8 - nbits * (key + 1);
let segment = (byte as usize >> shift) & ((1 << nbits) - 1);
let mut rev = 0usize;
for b in 0..nbits {
if segment & (1 << b) != 0 {
rev |= 1 << (nbits - 1 - b);
}
}
rev
}
}
#[cfg(test)]
mod tests {
use super::*;
fn encoder_pack(buckets: &[usize], nbits: usize) -> Vec<u8> {
let mut out = vec![0u8; buckets.len() * nbits / 8];
let mut bit_idx = 0usize;
for &bucket in buckets {
for b in 0..nbits {
let bit = ((bucket >> b) & 1) as u8;
out[bit_idx / 8] |= bit << (7 - (bit_idx % 8));
bit_idx += 1;
}
}
out
}
#[test]
fn colbert_packing_matches_encoder_bit_loop() {
for nbits in [1usize, 2, 4, 8] {
let p = ColbertPacking::new(nbits).unwrap();
let kpb = p.keys_per_byte();
let n = 1usize << nbits;
let mut buckets: Vec<usize> = Vec::new();
for b in 0..n {
for _ in 0..kpb {
buckets.push(b);
}
}
let mut x = 0x9E3779B9u32;
for _ in 0..(64 * kpb) {
x ^= x << 13;
x ^= x >> 17;
x ^= x << 5;
buckets.push(x as usize % n);
}
let bytes = encoder_pack(&buckets, nbits);
for (d, &want) in buckets.iter().enumerate() {
let got = p.bucket_index(bytes[d / kpb], d % kpb);
assert_eq!(got, want, "nbits={nbits} dim={d}");
}
let mut repacked = vec![0u8; bytes.len()];
p.pack_row(&buckets, &mut repacked);
assert_eq!(repacked, bytes, "nbits={nbits}");
}
}
}