pub const DEFAULT_BITS: usize = 2;
pub const MAX_LEVELS: usize = 16;
const ROUNDS: usize = 1;
const SPAN: f32 = 2.5;
pub fn code_len(padded: usize, bits: usize) -> usize { padded * bits / 8 }
pub fn pad_dim(dim: usize) -> usize {
dim.next_power_of_two().max(64)
}
fn signs(seed: u64, round: usize, padded: usize) -> Vec<f32> {
let mut s = seed ^ (round as u64).wrapping_mul(0x9E37_79B9_7F4A_7C15);
(0..padded).map(|_| {
s ^= s << 13; s ^= s >> 7; s ^= s << 17;
if s & 1 == 1 { 1.0 } else { -1.0 }
}).collect()
}
fn fwht(v: &mut [f32]) {
let n = v.len();
let mut h = 1;
while h < n {
for block in v.chunks_exact_mut(2 * h) {
let (a, b) = block.split_at_mut(h);
for j in 0..h {
let (x, y) = (a[j], b[j]);
a[j] = x + y;
b[j] = x - y;
}
}
h *= 2;
}
}
pub struct Encoder {
padded: usize,
bits: usize,
levels: usize,
signs: Vec<Vec<f32>>,
scale: f32,
table: [f32; MAX_LEVELS],
}
impl Encoder {
pub fn new(dim: usize, seed: u64, bits: usize) -> Encoder {
let padded = pad_dim(dim);
Encoder {
padded,
bits,
levels: 1 << bits,
signs: (0..ROUNDS).map(|r| signs(seed, r, padded)).collect(),
scale: (1.0 / (padded as f32).sqrt()).powi(ROUNDS as i32),
table: level_table(padded, bits),
}
}
pub fn bits(&self) -> usize { self.bits }
pub fn padded(&self) -> usize { self.padded }
pub fn rotate(&self, x: &mut [f32]) {
for sg in &self.signs {
for (xi, s) in x.iter_mut().zip(sg) { *xi *= s; }
fwht(x);
}
for xi in x.iter_mut() { *xi *= self.scale; }
}
pub fn encode(&self, v: &[f32]) -> (f32, Vec<u8>) {
let norm = v.iter().map(|a| a * a).sum::<f32>().sqrt();
let mut x = vec![0f32; self.padded];
if norm > 0.0 {
for (xi, vi) in x.iter_mut().zip(v) { *xi = vi / norm; }
}
self.rotate(&mut x);
let sd = (self.padded as f32).sqrt();
let per = 8 / self.bits;
let mut code = vec![0u8; self.padded * self.bits / 8];
for (j, &xi) in x.iter().enumerate() {
let z = (xi * sd).clamp(-SPAN, SPAN);
let q = (((z + SPAN) / (2.0 * SPAN)) * (self.levels as f32 - 1.0)).round() as usize;
let q = q.min(self.levels - 1) as u8;
code[j / per] |= q << ((j % per) * self.bits);
}
(norm, code)
}
pub fn rotate_query(&self, q: &[f32]) -> Vec<f32> {
let mut x = vec![0f32; self.padded];
x[..q.len()].copy_from_slice(q);
self.rotate(&mut x);
x
}
}
pub struct AffineQuery {
pub qrot: Vec<f32>,
pub a: f32,
pub b: f32,
}
impl Encoder {
pub fn affine_query(&self, q: &[f32]) -> AffineQuery {
let qrot = self.rotate_query(q);
let sd = (self.padded as f32).sqrt();
let step = 2.0 * SPAN / (self.levels as f32 - 1.0);
let sum_q: f32 = qrot.iter().sum();
AffineQuery { qrot, a: step / sd, b: -SPAN / sd * sum_q }
}
}
static LANE4: [[f32; 4]; 256] = {
let mut t = [[0f32; 4]; 256];
let mut b = 0usize;
while b < 256 {
t[b] = [(b & 3) as f32, ((b >> 2) & 3) as f32,
((b >> 4) & 3) as f32, ((b >> 6) & 3) as f32];
b += 1;
}
t
};
pub fn dot_est_affine2(norm: f32, code: &[u8], aq: &AffineQuery) -> f32 {
let q = &aq.qrot;
let (mut a0, mut a1, mut a2, mut a3) = (0f32, 0f32, 0f32, 0f32);
let cb = code.chunks_exact(16);
let rest = cb.remainder();
let qb = q.chunks_exact(64);
for (ch, qk) in cb.zip(qb) {
for i in 0..16 {
let l = &LANE4[ch[i] as usize];
let base = i * 4;
a0 += l[0] * qk[base];
a1 += l[1] * qk[base + 1];
a2 += l[2] * qk[base + 2];
a3 += l[3] * qk[base + 3];
}
}
let mut j = (code.len() - rest.len()) * 4;
for &b in rest {
a0 += (b & 3) as f32 * q[j];
a1 += ((b >> 2) & 3) as f32 * q[j + 1];
a2 += ((b >> 4) & 3) as f32 * q[j + 2];
a3 += ((b >> 6) & 3) as f32 * q[j + 3];
j += 4;
}
let s = (a0 + a1) + (a2 + a3);
norm * (aq.a * s + aq.b)
}
pub fn dot_est_affine(norm: f32, code: &[u8], aq: &AffineQuery) -> f32 {
let q = &aq.qrot;
let (mut a0, mut a1, mut a2, mut a3) = (0f32, 0f32, 0f32, 0f32);
let chunks = code.chunks_exact(4);
let rest = chunks.remainder();
let mut j = 0usize;
for ch in chunks {
a0 += (ch[0] & 0x0F) as f32 * q[j] + (ch[0] >> 4) as f32 * q[j + 1];
a1 += (ch[1] & 0x0F) as f32 * q[j + 2] + (ch[1] >> 4) as f32 * q[j + 3];
a2 += (ch[2] & 0x0F) as f32 * q[j + 4] + (ch[2] >> 4) as f32 * q[j + 5];
a3 += (ch[3] & 0x0F) as f32 * q[j + 6] + (ch[3] >> 4) as f32 * q[j + 7];
j += 8;
}
for &b in rest {
a0 += (b & 0x0F) as f32 * q[j] + (b >> 4) as f32 * q[j + 1];
j += 2;
}
let s = (a0 + a1) + (a2 + a3);
norm * (aq.a * s + aq.b)
}
pub fn level_table(padded: usize, bits: usize) -> [f32; MAX_LEVELS] {
let sd = (padded as f32).sqrt();
let levels = 1 << bits;
let mut t = [0f32; MAX_LEVELS];
for l in 0..levels {
let z = (l as f32) / (levels as f32 - 1.0) * (2.0 * SPAN) - SPAN;
t[l] = z / sd;
}
t
}