use crate::rotate::Rotation;
fn words_of(dim: usize) -> usize {
dim.div_ceil(64)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Bits {
One,
Four,
}
impl Bits {
#[must_use]
pub fn count(self) -> usize {
match self {
Bits::One => 1,
Bits::Four => 4,
}
}
#[must_use]
pub fn query_bits(self) -> usize {
match self {
Bits::One => 4,
Bits::Four => 8,
}
}
fn top(self) -> u64 {
(1 << self.count()) - 1
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct Coded {
pub norm: f32,
pub scale: f32,
pub lo: f32,
pub delta: f32,
}
pub struct Quantizer {
rot: Rotation,
bits: Bits,
}
impl Quantizer {
#[must_use]
pub fn new(dim: usize, bits: Bits, seed: u64) -> Quantizer {
Quantizer {
rot: Rotation::new(dim, seed),
bits,
}
}
#[must_use]
pub fn dim(&self) -> usize {
self.rot.dim()
}
#[must_use]
pub fn bits(&self) -> Bits {
self.bits
}
#[must_use]
pub fn seed(&self) -> u64 {
self.rot.seed()
}
#[must_use]
pub fn code_bytes(&self) -> usize {
words_of(self.dim()) * 8 * self.bits.count()
}
#[must_use]
pub fn rotate(&self, v: &[f32]) -> Vec<f32> {
let mut x = v.to_vec();
self.rot.apply(&mut x);
x
}
pub fn encode(&self, v: &[f32], centroid: &[f32], code: &mut [u8]) -> Coded {
let mut x = self.residual(v, centroid);
let norm = length(&x);
if norm > 0.0 {
let by = 1.0 / norm;
for c in &mut x {
*c *= by;
}
}
self.rot.apply(&mut x);
self.write(&x, norm, code)
}
pub fn encode_rotated(&self, x: &[f32], centroid: &[f32], code: &mut [u8]) -> Coded {
let mut r = self.residual(x, centroid);
let norm = length(&r);
if norm > 0.0 {
let by = 1.0 / norm;
for c in &mut r {
*c *= by;
}
}
self.write(&r, norm, code)
}
fn write(&self, x: &[f32], norm: f32, code: &mut [u8]) -> Coded {
assert_eq!(
code.len(),
self.code_bytes(),
"a code here is {} bytes and the buffer is {}",
self.code_bytes(),
code.len()
);
code.fill(0);
if norm == 0.0 {
return Coded {
norm: 0.0,
scale: 1.0,
lo: 0.0,
delta: 0.0,
};
}
let mut coded = match self.bits {
Bits::One => sign_code(x, code),
Bits::Four => level_code(x, code),
};
coded.norm = norm;
coded
}
#[must_use]
pub fn query(&self, q: &[f32], centroid: &[f32]) -> Query {
let mut x = self.residual(q, centroid);
let norm = length(&x);
if norm > 0.0 {
let by = 1.0 / norm;
for c in &mut x {
*c *= by;
}
}
self.rot.apply(&mut x);
self.prepare(x, norm)
}
#[must_use]
pub fn query_rotated(&self, q: &[f32], centroid: &[f32]) -> Query {
let mut x = self.residual(q, centroid);
let norm = length(&x);
if norm > 0.0 {
let by = 1.0 / norm;
for c in &mut x {
*c *= by;
}
}
self.prepare(x, norm)
}
fn prepare(&self, x: Vec<f32>, norm: f32) -> Query {
let sum = x.iter().sum();
let words = words_of(self.dim());
let wide = self.bits.query_bits();
let top = (1u64 << wide) - 1;
let (lo, hi) = span(&x);
let delta = step(lo, hi, top);
let by = 1.0 / delta;
let mut planes = vec![0u64; wide * words];
for (i, &c) in x.iter().enumerate() {
let level = level_of(c, lo, by, top);
for (b, plane) in planes.chunks_exact_mut(words).enumerate() {
plane[i / 64] |= ((level >> b) & 1) << (i % 64);
}
}
Query {
bits: self.bits,
words,
rotated: x,
planes,
lo,
delta,
sum,
norm,
}
}
fn residual(&self, v: &[f32], centroid: &[f32]) -> Vec<f32> {
assert_eq!(
v.len(),
self.dim(),
"this collection holds {} dimensional vectors and was handed {}",
self.dim(),
v.len()
);
assert_eq!(
centroid.len(),
self.dim(),
"the centroid is {} dimensional and the collection is {}",
centroid.len(),
self.dim()
);
v.iter().zip(centroid).map(|(a, b)| a - b).collect()
}
}
pub struct Query {
bits: Bits,
words: usize,
rotated: Vec<f32>,
planes: Vec<u64>,
lo: f32,
delta: f32,
sum: f32,
norm: f32,
}
impl Query {
pub fn scan(&self, codes: &[u8], meta: &[Coded], out: &mut [f32]) {
assert_eq!(
meta.len(),
out.len(),
"{} codes and room for {} answers",
meta.len(),
out.len()
);
let stride = self.bits.count() * self.words * 8;
assert_eq!(
codes.len(),
meta.len() * stride,
"{} members of {stride} bytes is {} and there are {}",
meta.len(),
meta.len() * stride,
codes.len()
);
match self.words {
1 => self.scan_at::<1>(codes, meta, out),
2 => self.scan_at::<2>(codes, meta, out),
4 => self.scan_at::<4>(codes, meta, out),
6 => self.scan_at::<6>(codes, meta, out),
8 => self.scan_at::<8>(codes, meta, out),
12 => self.scan_at::<12>(codes, meta, out),
16 => self.scan_at::<16>(codes, meta, out),
24 => self.scan_at::<24>(codes, meta, out),
48 => self.scan_at::<48>(codes, meta, out),
_ => {
for (i, (code, coded)) in codes.chunks_exact(stride).zip(meta).enumerate() {
let (total, cross) =
packed_dot(code, self.bits.count(), self.words, &self.planes);
out[i] = self.settle(total, cross, coded);
}
}
}
}
fn scan_at<const W: usize>(&self, codes: &[u8], meta: &[Coded], out: &mut [f32]) {
let planes = self.bits.count();
let wide = self.bits.query_bits();
for (i, (code, coded)) in codes.chunks_exact(planes * W * 8).zip(meta).enumerate() {
let (total, cross) = fixed_dot::<W>(code, planes, wide, &self.planes);
out[i] = self.settle(total, cross, coded);
}
}
#[inline]
fn settle(&self, total: u32, cross: u32, coded: &Coded) -> f32 {
let levels = self.lo * total as f32 + self.delta * cross as f32;
let cos = (coded.lo * self.sum + coded.delta * levels) * coded.scale;
(self.norm * self.norm + coded.norm * coded.norm - 2.0 * self.norm * coded.norm * cos)
.max(0.0)
}
#[must_use]
pub fn distance(&self, code: &[u8], coded: &Coded) -> f32 {
let mut out = [0.0f32];
self.scan(code, std::slice::from_ref(coded), &mut out);
out[0]
}
#[must_use]
pub fn cosine(&self, code: &[u8], coded: &Coded) -> f32 {
if coded.norm == 0.0 || self.norm == 0.0 {
return 0.0;
}
let (total, cross) = packed_dot(code, self.bits.count(), self.words, &self.planes);
let levels = self.lo * total as f32 + self.delta * cross as f32;
(coded.lo * self.sum + coded.delta * levels) * coded.scale
}
#[must_use]
pub fn cosine_exact(&self, code: &[u8], coded: &Coded) -> f32 {
if coded.norm == 0.0 || self.norm == 0.0 {
return 0.0;
}
let levels = exact_dot(code, self.bits.count(), self.words, &self.rotated);
(coded.lo * self.sum + coded.delta * levels) * coded.scale
}
}
fn span(x: &[f32]) -> (f32, f32) {
let mut lo = f32::INFINITY;
let mut hi = f32::NEG_INFINITY;
for &c in x {
lo = lo.min(c);
hi = hi.max(c);
}
(lo, hi)
}
fn step(lo: f32, hi: f32, top: u64) -> f32 {
if hi > lo { (hi - lo) / top as f32 } else { 1.0 }
}
fn level_of(c: f32, lo: f32, by: f32, top: u64) -> u64 {
(((c - lo) * by).round() as i64).clamp(0, top as i64) as u64
}
fn put(code: &mut [u8], words: usize, plane: usize, w: usize, v: u64) {
let at = (plane * words + w) * 8;
code[at..at + 8].copy_from_slice(&v.to_le_bytes());
}
fn sign_code(x: &[f32], code: &mut [u8]) -> Coded {
let words = words_of(x.len());
let mut abs = 0.0f32;
for (w, chunk) in x.chunks(64).enumerate() {
let mut bits = 0u64;
for (k, &c) in chunk.iter().enumerate() {
abs += c.abs();
if c >= 0.0 {
bits |= 1 << k;
}
}
put(code, words, 0, w, bits);
}
let root = (x.len() as f32).sqrt();
Coded {
norm: 0.0,
scale: recip(abs / root),
lo: -1.0 / root,
delta: 2.0 / root,
}
}
fn level_code(x: &[f32], code: &mut [u8]) -> Coded {
let words = words_of(x.len());
let top = Bits::Four.top();
let (lo, hi) = span(x);
let delta = step(lo, hi, top);
let by = 1.0 / delta;
let mut recon = 0.0f32;
let mut dot = 0.0f32;
for (w, chunk) in x.chunks(64).enumerate() {
let mut planes = [0u64; 4];
for (k, &c) in chunk.iter().enumerate() {
let level = level_of(c, lo, by, top);
for (b, plane) in planes.iter_mut().enumerate() {
*plane |= ((level >> b) & 1) << k;
}
let back = lo + level as f32 * delta;
recon += back * back;
dot += back * c;
}
for (b, &plane) in planes.iter().enumerate() {
put(code, words, b, w, plane);
}
}
let len = recon.sqrt();
Coded {
norm: 0.0,
scale: recip(dot / len),
lo: lo / len,
delta: delta / len,
}
}
#[allow(clippy::chunks_exact_to_as_chunks)]
fn packed_dot(code: &[u8], planes: usize, words: usize, query: &[u64]) -> (u32, u32) {
assert_eq!(
code.len(),
planes * words * 8,
"a code here is {} bytes and this one is {}",
planes * words * 8,
code.len()
);
let mut total = 0u32;
let mut cross = 0u32;
for (a, plane) in code.chunks_exact(words * 8).enumerate() {
let mut ones = 0u32;
for chunk in plane.chunks_exact(8) {
ones += word(chunk).count_ones();
}
total += ones << a;
for (b, qp) in query.chunks_exact(words).enumerate() {
let mut acc = 0u32;
for (chunk, &qw) in plane.chunks_exact(8).zip(qp) {
acc += (word(chunk) & qw).count_ones();
}
cross += acc << (a + b);
}
}
(total, cross)
}
#[inline]
fn fixed_dot<const W: usize>(code: &[u8], planes: usize, wide: usize, query: &[u64]) -> (u32, u32) {
let mut total = 0u32;
let mut cross = 0u32;
for a in 0..planes {
let mut plane = [0u64; W];
let at = a * W * 8;
let mut ones = 0u32;
for (w, slot) in plane.iter_mut().enumerate() {
*slot = word(&code[at + w * 8..at + w * 8 + 8]);
ones += slot.count_ones();
}
total += ones << a;
for b in 0..wide {
let qp = &query[b * W..(b + 1) * W];
let mut acc = 0u32;
for (w, &c) in plane.iter().enumerate() {
acc += (c & qp[w]).count_ones();
}
cross += acc << (a + b);
}
}
(total, cross)
}
fn recip(x: f32) -> f32 {
if x == 0.0 { 0.0 } else { 1.0 / x }
}
fn exact_dot(code: &[u8], planes: usize, words: usize, query: &[f32]) -> f32 {
assert_eq!(
code.len(),
planes * words * 8,
"a code here is {} bytes and this one is {}",
planes * words * 8,
code.len()
);
let mut sum = 0.0f32;
for (i, &qi) in query.iter().enumerate() {
let mut level = 0u32;
for b in 0..planes {
let at = (b * words + i / 64) * 8;
level |= (((word(&code[at..at + 8]) >> (i % 64)) & 1) as u32) << b;
}
sum += level as f32 * qi;
}
sum
}
fn word(bytes: &[u8]) -> u64 {
u64::from_le_bytes(bytes.try_into().expect("eight bytes"))
}
fn length(v: &[f32]) -> f32 {
v.iter().map(|c| c * c).sum::<f32>().sqrt()
}
#[cfg(test)]
mod tests {
use super::*;
use yo_common::Rng;
fn corpus(dim: usize, n: usize, seed: u64) -> Vec<Vec<f32>> {
let mut rng = Rng::new(seed);
(0..n)
.map(|_| {
let mut v: Vec<f32> = (0..dim)
.map(|i| {
let u = (rng.next_u64() >> 40) as f32 / (1u32 << 24) as f32;
let heavy = if i < dim / 16 { 6.0 } else { 1.0 };
(u * 2.0 - 1.0) * heavy
})
.collect();
let len = length(&v);
for c in &mut v {
*c /= len;
}
v
})
.collect()
}
fn mean(vs: &[Vec<f32>]) -> Vec<f32> {
let dim = vs[0].len();
let mut c = vec![0.0f32; dim];
for v in vs {
for (a, b) in c.iter_mut().zip(v) {
*a += b;
}
}
for a in &mut c {
*a /= vs.len() as f32;
}
c
}
fn exact(a: &[f32], b: &[f32]) -> f32 {
a.iter().zip(b).map(|(x, y)| (x - y) * (x - y)).sum()
}
fn residual(v: &[f32], c: &[f32]) -> Vec<f32> {
v.iter().zip(c).map(|(a, b)| a - b).collect()
}
fn dot(a: &[f32], b: &[f32]) -> f32 {
a.iter().zip(b).map(|(x, y)| x * y).sum()
}
fn encode_all(q: &Quantizer, vs: &[Vec<f32>], c: &[f32]) -> (Vec<u8>, Vec<Coded>) {
let width = q.code_bytes();
let mut codes = vec![0u8; width * vs.len()];
let mut meta = Vec::with_capacity(vs.len());
for (i, v) in vs.iter().enumerate() {
meta.push(q.encode(v, c, &mut codes[i * width..(i + 1) * width]));
}
(codes, meta)
}
fn recall(bits: Bits, dim: usize, keep: usize) -> f32 {
let vs = corpus(dim, 800, 1);
let qs = corpus(dim, 40, 2);
let c = mean(&vs);
let q = Quantizer::new(dim, bits, 99);
let (codes, meta) = encode_all(&q, &vs, &c);
let width = q.code_bytes();
let mut hits = 0usize;
for query in &qs {
let mut truth: Vec<(usize, f32)> = vs
.iter()
.enumerate()
.map(|(i, v)| (i, exact(query, v)))
.collect();
truth.sort_by(|a, b| a.1.total_cmp(&b.1));
let want: Vec<usize> = truth[..10].iter().map(|(i, _)| *i).collect();
let prepared = q.query(query, &c);
let mut guess: Vec<(usize, f32)> = (0..vs.len())
.map(|i| {
let code = &codes[i * width..(i + 1) * width];
(i, prepared.distance(code, &meta[i]))
})
.collect();
guess.sort_by(|a, b| a.1.total_cmp(&b.1));
let got: Vec<usize> = guess[..keep].iter().map(|(i, _)| *i).collect();
hits += want.iter().filter(|i| got.contains(i)).count();
}
hits as f32 / (qs.len() * 10) as f32
}
#[test]
#[ignore = "prints a table rather than asserting anything"]
fn recall_table() {
for (bits, name) in [(Bits::One, "1 bit"), (Bits::Four, "4 bit")] {
for dim in [128usize, 256, 768] {
for keep in [10usize, 20, 40, 100] {
println!(
"{name} dim {dim} keep {keep}: {:.3}",
recall(bits, dim, keep)
);
}
}
}
}
#[test]
fn a_code_is_the_width_it_says_it_is() {
assert_eq!(Quantizer::new(768, Bits::One, 1).code_bytes(), 96);
assert_eq!(Quantizer::new(768, Bits::Four, 1).code_bytes(), 384);
assert_eq!(Quantizer::new(100, Bits::One, 1).code_bytes(), 16);
assert_eq!(Quantizer::new(128, Bits::One, 1).code_bytes(), 16);
}
#[test]
fn one_bit_finds_the_true_neighbours_inside_a_short_rerank() {
let r = recall(Bits::One, 256, 40);
assert!(r >= 0.95, "recall at 10 was {r}");
}
#[test]
fn four_bits_is_better_than_one() {
let one = recall(Bits::One, 128, 20);
let four = recall(Bits::Four, 128, 20);
assert!(four > one, "one bit got {one} and four bits got {four}");
assert!(four >= 0.95, "four bit recall at 10 was {four}");
}
#[test]
fn the_estimate_is_close_to_the_truth_rather_than_merely_ordered() {
let dim = 256;
let vs = corpus(dim, 200, 5);
let qs = corpus(dim, 20, 6);
let c = mean(&vs);
let q = Quantizer::new(dim, Bits::One, 3);
let (codes, meta) = encode_all(&q, &vs, &c);
let width = q.code_bytes();
let mut worst = 0.0f32;
let mut bias = 0.0f32;
let mut n = 0usize;
for query in &qs {
let prepared = q.query(query, &c);
for (i, v) in vs.iter().enumerate() {
let truth = exact(query, v);
let guess = prepared.distance(&codes[i * width..(i + 1) * width], &meta[i]);
let err = (guess - truth) / truth;
worst = worst.max(err.abs());
bias += err;
n += 1;
}
}
let bias = bias / n as f32;
assert!(
bias.abs() < 0.02,
"the estimate is off by {bias} on average"
);
assert!(worst < 0.5, "the worst estimate was off by {worst}");
}
#[test]
fn the_query_is_quantised_finer_than_the_code_it_is_measured_against() {
for (bits, dim) in [
(Bits::One, 128),
(Bits::One, 256),
(Bits::One, 768),
(Bits::Four, 256),
(Bits::Four, 768),
] {
let vs = corpus(dim, 200, 5);
let qs = corpus(dim, 20, 6);
let c = mean(&vs);
let q = Quantizer::new(dim, bits, 3);
let (codes, meta) = encode_all(&q, &vs, &c);
let width = q.code_bytes();
let mut from_query = 0.0f32;
let mut from_code = 0.0f32;
for query in &qs {
let prepared = q.query(query, &c);
let qr = residual(query, &c);
let qn = length(&qr);
for (i, v) in vs.iter().enumerate() {
let code = &codes[i * width..(i + 1) * width];
let fast = prepared.cosine(code, &meta[i]);
let slow = prepared.cosine_exact(code, &meta[i]);
let vr = residual(v, &c);
let truth = dot(&qr, &vr) / (qn * length(&vr));
from_query += (fast - slow).abs();
from_code += (slow - truth).abs();
}
}
let ratio = from_query / from_code;
assert!(
ratio < 0.5,
"{dim} at {bits:?}: the query costs {ratio} of what the code costs"
);
}
}
#[test]
fn a_vector_sitting_on_its_centroid_is_not_a_division_by_zero() {
let q = Quantizer::new(16, Bits::One, 1);
let c = vec![0.5f32; 16];
let mut code = vec![0u8; q.code_bytes()];
let coded = q.encode(&c, &c, &mut code);
assert_eq!(coded.norm, 0.0);
assert!(code.iter().all(|b| *b == 0));
let query = q.query(&[1.0f32; 16], &c);
let d = query.distance(&code, &coded);
assert!(d.is_finite(), "{d}");
let want: f32 = (0..16).map(|_| 0.25f32).sum();
assert!((d - want).abs() < 1e-3, "{d} against {want}");
}
#[test]
fn a_query_sitting_on_the_centroid_is_not_a_division_by_zero() {
let q = Quantizer::new(16, Bits::One, 1);
let c = vec![0.5f32; 16];
let mut code = vec![0u8; q.code_bytes()];
let v: Vec<f32> = (0..16).map(|i| i as f32 * 0.1).collect();
let coded = q.encode(&v, &c, &mut code);
let d = q.query(&c, &c).distance(&code, &coded);
assert!(d.is_finite(), "{d}");
}
#[test]
fn a_code_is_written_over_whatever_was_in_the_buffer() {
let q = Quantizer::new(32, Bits::One, 1);
let c = vec![0.0f32; 32];
let v: Vec<f32> = (0..32).map(|i| (i as f32).sin()).collect();
let mut fresh = vec![0u8; q.code_bytes()];
let mut dirty = vec![0xffu8; q.code_bytes()];
let a = q.encode(&v, &c, &mut fresh);
let b = q.encode(&v, &c, &mut dirty);
assert_eq!(fresh, dirty);
assert_eq!(a, b);
assert!(fresh[4..].iter().all(|b| *b == 0), "{fresh:?}");
}
#[test]
fn the_same_seed_is_the_same_code() {
let v: Vec<f32> = (0..64).map(|i| (i as f32 * 0.3).cos()).collect();
let c = vec![0.0f32; 64];
let mut a = vec![0u8; 8];
let mut b = vec![0u8; 8];
Quantizer::new(64, Bits::One, 12).encode(&v, &c, &mut a);
Quantizer::new(64, Bits::One, 12).encode(&v, &c, &mut b);
assert_eq!(a, b);
let mut d = vec![0u8; 8];
Quantizer::new(64, Bits::One, 13).encode(&v, &c, &mut d);
assert_ne!(a, d, "two seeds should not be one code");
}
#[test]
fn a_code_reads_back_as_the_levels_it_was_written_from() {
let dim = 200;
let q = Quantizer::new(dim, Bits::Four, 4);
let c = vec![0.0f32; dim];
let v: Vec<f32> = (0..dim).map(|i| (i as f32 * 0.11).sin()).collect();
let mut code = vec![0u8; q.code_bytes()];
let coded = q.encode(&v, &c, &mut code);
let words = words_of(dim);
let mut back = vec![0.0f32; dim];
for (i, b) in back.iter_mut().enumerate() {
let mut level = 0u32;
for p in 0..4 {
let at = (p * words + i / 64) * 8;
level |= (((word(&code[at..at + 8]) >> (i % 64)) & 1) as u32) << p;
}
*b = coded.lo + level as f32 * coded.delta;
}
let mut spun: Vec<f32> = v.iter().map(|c| c / length(&v)).collect();
crate::Rotation::new(dim, q.seed()).apply(&mut spun);
let cos = back
.iter()
.zip(&spun)
.map(|(a, b)| a * b)
.sum::<f32>()
.abs()
/ length(&back);
assert!(cos > 0.95, "the code points somewhere else: {cos}");
}
}