use crate::kernel::{self, Kernel};
use crate::packing::Packing;
use crate::Error;
pub const MAX_DIM: usize = 256;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct NibbleTables {
pub tables: [[i8; 16]; 8],
pub from_hi: [bool; 8],
}
#[derive(Debug, Clone)]
pub struct Lut {
fused: Vec<i8>,
keys_per_byte: usize,
nbits: usize,
scale: f32,
nibble: Option<NibbleTables>,
force_scalar: bool,
pin: Option<Kernel>,
}
impl Lut {
pub fn new<P: Packing>(packing: &P, bucket_weights: &[f32]) -> Result<Self, Error> {
let nbits = packing.nbits();
if !matches!(nbits, 1 | 2 | 4 | 8) {
return Err(Error::NbitsUnsupported(nbits));
}
let expected = 1usize << nbits;
if bucket_weights.len() != expected {
return Err(Error::BucketCount {
expected,
got: bucket_weights.len(),
});
}
let max_abs = bucket_weights.iter().fold(0.0f32, |m, &x| m.max(x.abs()));
let scale = (max_abs / 127.0).max(1e-12);
let vals: Vec<i8> = bucket_weights
.iter()
.map(|&w| (w / scale).round().clamp(-127.0, 127.0) as i8)
.collect();
let keys_per_byte = 8 / nbits;
let mut fused = vec![0i8; 256 * keys_per_byte];
for byte in 0..256usize {
for k in 0..keys_per_byte {
let bi = packing.bucket_index(byte as u8, k);
assert!(
bi < expected,
"Packing::bucket_index returned {bi} >= 2^nbits for byte {byte} key {k}"
);
fused[byte * keys_per_byte + k] = vals[bi];
}
}
let nibble = derive_nibble_tables(&fused, keys_per_byte);
Ok(Self {
fused,
keys_per_byte,
nbits,
scale,
nibble,
force_scalar: false,
pin: None,
})
}
pub fn colbert(nbits: usize, bucket_weights: &[f32]) -> Result<Self, Error> {
Self::new(&crate::ColbertPacking::new(nbits)?, bucket_weights)
}
pub fn force_scalar(mut self, yes: bool) -> Self {
self.force_scalar = yes;
self
}
pub fn pin_kernel(mut self, kernel: Option<Kernel>) -> Self {
self.pin = kernel;
self
}
pub fn nbits(&self) -> usize {
self.nbits
}
pub fn keys_per_byte(&self) -> usize {
self.keys_per_byte
}
pub fn scale(&self) -> f32 {
self.scale
}
#[inline]
pub fn expand(&self, byte: u8) -> &[i8] {
let base = byte as usize * self.keys_per_byte;
&self.fused[base..base + self.keys_per_byte]
}
pub fn fused_table(&self) -> &[i8] {
&self.fused
}
pub fn nibble_tables(&self) -> Option<&NibbleTables> {
self.nibble.as_ref()
}
pub fn kernel(&self, dim: usize) -> Kernel {
kernel::select(self, dim)
}
pub(crate) fn force_scalar_set(&self) -> bool {
self.force_scalar
}
pub(crate) fn pinned_kernel(&self) -> Option<Kernel> {
self.pin
}
pub(crate) fn fingerprint(&self) -> (usize, u32) {
(self.nbits, self.scale.to_bits())
}
}
fn derive_nibble_tables(fused: &[i8], keys_per_byte: usize) -> Option<NibbleTables> {
if keys_per_byte > 8 || keys_per_byte == 1 {
return None;
}
let mut tables = [[0i8; 16]; 8];
let mut from_hi = [false; 8];
for k in 0..keys_per_byte {
let hi: [i8; 16] = std::array::from_fn(|x| fused[(x << 4) * keys_per_byte + k]);
if (0..256).all(|b| fused[b * keys_per_byte + k] == hi[b >> 4]) {
tables[k] = hi;
from_hi[k] = true;
continue;
}
let lo: [i8; 16] = std::array::from_fn(|x| fused[x * keys_per_byte + k]);
if (0..256).all(|b| fused[b * keys_per_byte + k] == lo[b & 15]) {
tables[k] = lo;
from_hi[k] = false;
continue;
}
return None;
}
Some(NibbleTables { tables, from_hi })
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ColbertPacking;
fn weights(nbits: usize) -> Vec<f32> {
let n = 1usize << nbits;
(0..n)
.map(|i| -0.35 + 0.7 * (i as f32 + 0.5) / n as f32)
.collect()
}
#[test]
fn fused_table_expands_to_quantised_bucket_weights() {
for nbits in [1usize, 2, 4, 8] {
let p = ColbertPacking::new(nbits).unwrap();
let w = weights(nbits);
let lut = Lut::new(&p, &w).unwrap();
for byte in 0..=255u8 {
for k in 0..lut.keys_per_byte() {
let bi = p.bucket_index(byte, k);
let want = (w[bi] / lut.scale()).round().clamp(-127.0, 127.0) as i8;
assert_eq!(lut.expand(byte)[k], want, "nbits={nbits} byte={byte} k={k}");
}
}
}
}
#[test]
fn nibble_factorisation_holds_for_sub_byte_codes() {
for nbits in [1usize, 2, 4] {
let lut = Lut::colbert(nbits, &weights(nbits)).unwrap();
let nib = lut
.nibble_tables()
.unwrap_or_else(|| panic!("nbits={nbits}: not nibble-separable"));
for b in 0..256usize {
for k in 0..lut.keys_per_byte() {
let nibble = if nib.from_hi[k] { b >> 4 } else { b & 15 };
assert_eq!(
lut.fused_table()[b * lut.keys_per_byte() + k],
nib.tables[k][nibble]
);
}
}
}
assert!(Lut::colbert(8, &weights(8)).unwrap().nibble_tables().is_none());
}
#[test]
fn rejects_bad_inputs() {
assert_eq!(ColbertPacking::new(3).unwrap_err(), Error::NbitsUnsupported(3));
assert_eq!(
Lut::colbert(4, &[0.0; 15]).unwrap_err(),
Error::BucketCount {
expected: 16,
got: 15
}
);
}
}