use rand::{rngs::StdRng, RngExt, SeedableRng};
use serde::{Deserialize, Serialize};
const DEFAULT_SEED: u64 = 0x5241_4249_5451_5121;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RaBitQuantizer {
dim: usize,
centroid: Vec<f32>,
rotation: Vec<f32>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RaBitCode {
pub bits: Vec<u8>,
pub dtc_sq: f32,
pub est_factor: f32,
}
pub struct PreparedQuery {
rq: Vec<f32>,
qn_sq: f32,
}
impl PreparedQuery {
pub fn rq(&self) -> &[f32] {
&self.rq
}
pub fn qn_sq(&self) -> f32 {
self.qn_sq
}
}
impl RaBitQuantizer {
pub fn fit(training_vectors: &[Vec<f32>]) -> Self {
Self::fit_with_seed(training_vectors, DEFAULT_SEED)
}
pub fn fit_with_seed(training_vectors: &[Vec<f32>], seed: u64) -> Self {
assert!(
!training_vectors.is_empty(),
"Need at least one training vector"
);
let dim = training_vectors[0].len();
assert!(dim > 0, "Dimension must be positive");
let mut centroid = vec![0.0f32; dim];
for v in training_vectors {
assert_eq!(v.len(), dim, "Inconsistent vector dimensions");
for (c, &x) in centroid.iter_mut().zip(v.iter()) {
*c += x;
}
}
let inv_n = 1.0 / training_vectors.len() as f32;
for c in &mut centroid {
*c *= inv_n;
}
let rotation = random_orthonormal(dim, seed);
Self {
dim,
centroid,
rotation,
}
}
pub fn dim(&self) -> usize {
self.dim
}
pub fn encode(&self, vector: &[f32]) -> RaBitCode {
debug_assert_eq!(vector.len(), self.dim);
let mut res = vec![0.0f32; self.dim];
let mut dtc_sq = 0.0f32;
for ((r, &v), &c) in res.iter_mut().zip(vector).zip(&self.centroid) {
*r = v - c;
dtc_sq += (v - c) * (v - c);
}
let ro = self.matvec(&res);
let mut bits = vec![0u8; self.bytes()];
let mut l1 = 0.0f32;
for (i, &x) in ro.iter().enumerate() {
l1 += x.abs();
if x >= 0.0 {
bits[i / 8] |= 1 << (i % 8);
}
}
let est_factor = if l1 > f32::EPSILON { dtc_sq / l1 } else { 0.0 };
RaBitCode {
bits,
dtc_sq,
est_factor,
}
}
pub fn prepare_query(&self, query: &[f32]) -> PreparedQuery {
debug_assert_eq!(query.len(), self.dim);
let mut res = vec![0.0f32; self.dim];
let mut qn_sq = 0.0f32;
for ((r, &q), &c) in res.iter_mut().zip(query).zip(&self.centroid) {
*r = q - c;
qn_sq += (q - c) * (q - c);
}
let rq = self.matvec(&res);
PreparedQuery { rq, qn_sq }
}
pub fn estimate_dist_sq(&self, query: &PreparedQuery, code: &RaBitCode) -> f32 {
let mut s = 0.0f32;
for (i, &rq) in query.rq.iter().enumerate() {
let bit = (code.bits[i / 8] >> (i % 8)) & 1;
if bit == 1 {
s += rq;
} else {
s -= rq;
}
}
let dsq = code.dtc_sq + query.qn_sq - 2.0 * code.est_factor * s;
dsq.max(0.0)
}
fn bytes(&self) -> usize {
self.dim.div_ceil(8)
}
fn matvec(&self, v: &[f32]) -> Vec<f32> {
let d = self.dim;
let mut out = vec![0.0f32; d];
for (r, o) in out.iter_mut().enumerate() {
let row = &self.rotation[r * d..(r + 1) * d];
*o = super::simd::dot_product_simd(row, v);
}
out
}
fn matvec_transpose(&self, v: &[f32]) -> Vec<f32> {
let d = self.dim;
let mut out = vec![0.0f32; d];
for (r, &vr) in v.iter().enumerate() {
let row = &self.rotation[r * d..(r + 1) * d];
for (o, &rc) in out.iter_mut().zip(row) {
*o += rc * vr;
}
}
out
}
}
impl RaBitQuantizer {
pub fn quantize(&self, vector: &[f32]) -> RaBitCode {
self.encode(vector)
}
pub fn dequantize(&self, quantized: &RaBitCode) -> Vec<f32> {
let d = self.dim;
let inv_sqrt_d = 1.0 / (d as f32).sqrt();
let mut xbar = vec![0.0f32; d];
for (i, x) in xbar.iter_mut().enumerate() {
let bit = (quantized.bits[i / 8] >> (i % 8)) & 1;
*x = if bit == 1 { inv_sqrt_d } else { -inv_sqrt_d };
}
let dir = self.matvec_transpose(&xbar);
let dtc = quantized.dtc_sq.sqrt();
dir.iter()
.zip(&self.centroid)
.map(|(&u, &c)| c + dtc * u)
.collect()
}
pub fn distance_quantized(&self, a: &RaBitCode, b: &RaBitCode) -> f32 {
let a_full = self.dequantize(a);
self.distance_asymmetric(&a_full, b)
}
pub fn distance_asymmetric(&self, query: &[f32], quantized: &RaBitCode) -> f32 {
let prepared = self.prepare_query(query);
self.estimate_dist_sq(&prepared, quantized).sqrt()
}
}
fn random_orthonormal(dim: usize, seed: u64) -> Vec<f32> {
let mut rng = StdRng::seed_from_u64(seed);
let mut rows: Vec<Vec<f32>> = Vec::with_capacity(dim);
for _ in 0..dim {
let mut v: Vec<f32> = (0..dim).map(|_| gaussian(&mut rng)).collect();
for prev in &rows {
let proj = dot(&v, prev);
for (vi, &pi) in v.iter_mut().zip(prev) {
*vi -= proj * pi;
}
}
let mut norm = dot(&v, &v).sqrt();
while norm < 1e-6 {
v = (0..dim).map(|_| gaussian(&mut rng)).collect();
for prev in &rows {
let proj = dot(&v, prev);
for (vi, &pi) in v.iter_mut().zip(prev) {
*vi -= proj * pi;
}
}
norm = dot(&v, &v).sqrt();
}
let inv = 1.0 / norm;
for vi in &mut v {
*vi *= inv;
}
rows.push(v);
}
let mut flat = Vec::with_capacity(dim * dim);
for row in rows {
flat.extend_from_slice(&row);
}
flat
}
#[inline]
fn dot(a: &[f32], b: &[f32]) -> f32 {
a.iter().zip(b).map(|(&x, &y)| x * y).sum()
}
#[inline]
fn gaussian(rng: &mut StdRng) -> f32 {
let u1: f32 = rng.random::<f32>().max(1e-7);
let u2: f32 = rng.random::<f32>();
(-2.0 * u1.ln()).sqrt() * (std::f32::consts::TAU * u2).cos()
}
#[cfg(test)]
mod tests {
use super::*;
fn rng_vec(rng: &mut StdRng, dim: usize) -> Vec<f32> {
(0..dim).map(|_| rng.random::<f32>() * 2.0 - 1.0).collect()
}
#[test]
fn rotation_is_orthonormal() {
let d = 64;
let r = random_orthonormal(d, 123);
for i in 0..d {
for j in 0..d {
let ri = &r[i * d..(i + 1) * d];
let rj = &r[j * d..(j + 1) * d];
let prod = dot(ri, rj);
let expected = if i == j { 1.0 } else { 0.0 };
assert!(
(prod - expected).abs() < 1e-3,
"R·Rᵀ[{i},{j}] = {prod}, expected {expected}"
);
}
}
}
#[test]
fn rotation_is_deterministic() {
assert_eq!(random_orthonormal(32, 42), random_orthonormal(32, 42));
}
#[test]
fn estimator_is_approximately_unbiased() {
let mut rng = StdRng::seed_from_u64(7);
let dim = 128;
let train: Vec<Vec<f32>> = (0..500).map(|_| rng_vec(&mut rng, dim)).collect();
let q = RaBitQuantizer::fit(&train);
let mut rel_errs = Vec::new();
for _ in 0..200 {
let o = rng_vec(&mut rng, dim);
let query = rng_vec(&mut rng, dim);
let code = q.encode(&o);
let prep = q.prepare_query(&query);
let est = q.estimate_dist_sq(&prep, &code);
let truth: f32 = o.iter().zip(&query).map(|(a, b)| (a - b) * (a - b)).sum();
rel_errs.push((est - truth) / truth);
}
let mean_bias: f32 = rel_errs.iter().sum::<f32>() / rel_errs.len() as f32;
assert!(
mean_bias.abs() < 0.10,
"estimator mean relative bias too large: {mean_bias}"
);
}
#[test]
fn rerank_recall_beats_hamming_floor() {
let mut rng = StdRng::seed_from_u64(99);
let dim = 128;
let n = 2000;
let base: Vec<Vec<f32>> = (0..n).map(|_| rng_vec(&mut rng, dim)).collect();
let q = RaBitQuantizer::fit(&base);
let codes: Vec<RaBitCode> = base.iter().map(|v| q.encode(v)).collect();
let k = 10;
let rerank = 100; let mut total_recall = 0.0;
let trials = 50;
for _ in 0..trials {
let query = rng_vec(&mut rng, dim);
let mut exact: Vec<(f32, usize)> = base
.iter()
.enumerate()
.map(|(i, v)| {
(
v.iter().zip(&query).map(|(a, b)| (a - b) * (a - b)).sum(),
i,
)
})
.collect();
exact.sort_by(|a, b| a.0.total_cmp(&b.0));
let truth: std::collections::HashSet<usize> =
exact.iter().take(k).map(|(_, i)| *i).collect();
let prep = q.prepare_query(&query);
let mut est: Vec<(f32, usize)> = codes
.iter()
.enumerate()
.map(|(i, c)| (q.estimate_dist_sq(&prep, c), i))
.collect();
est.sort_by(|a, b| a.0.total_cmp(&b.0));
let mut pool: Vec<(f32, usize)> = est
.iter()
.take(rerank)
.map(|&(_, i)| {
let d: f32 = base[i]
.iter()
.zip(&query)
.map(|(a, b)| (a - b) * (a - b))
.sum();
(d, i)
})
.collect();
pool.sort_by(|a, b| a.0.total_cmp(&b.0));
let got: std::collections::HashSet<usize> =
pool.iter().take(k).map(|(_, i)| *i).collect();
total_recall += truth.intersection(&got).count() as f32 / k as f32;
}
let recall = total_recall / trials as f32;
assert!(recall > 0.80, "RaBitQ rerank recall@10 too low: {recall}");
}
#[test]
fn quantizer_trait_roundtrip() {
let mut rng = StdRng::seed_from_u64(5);
let dim = 96;
let train: Vec<Vec<f32>> = (0..200).map(|_| rng_vec(&mut rng, dim)).collect();
let q = RaBitQuantizer::fit(&train);
let v = rng_vec(&mut rng, dim);
let code = q.quantize(&v);
assert_eq!(code.bits.len(), dim.div_ceil(8));
let self_d = q.distance_asymmetric(&v, &code);
let other = rng_vec(&mut rng, dim);
let other_d = q.distance_asymmetric(&other, &code);
assert!(
self_d < other_d,
"self distance {self_d} should be < cross distance {other_d}"
);
assert_eq!(q.dequantize(&code).len(), dim);
}
}