use yo_common::Rng;
#[derive(Debug, Clone, Copy)]
pub struct Shape {
pub ksim: usize,
pub dproj: usize,
pub reps: usize,
}
impl Default for Shape {
fn default() -> Shape {
Shape {
ksim: 4,
dproj: 16,
reps: 8,
}
}
}
pub struct Encoder {
dim: usize,
shape: Shape,
planes: Vec<f32>,
proj: Vec<f32>,
}
impl Encoder {
#[must_use]
pub fn new(dim: usize, shape: Shape, seed: u64) -> Encoder {
assert!(dim > 0, "a token has to have a dimension");
assert!(
(1..=16).contains(&shape.ksim),
"ksim is {}, and a bucket index is built by shifting, so it has to \
stay somewhere a machine can count to",
shape.ksim
);
assert!(shape.dproj > 0, "a block has to have a width");
assert!(shape.reps > 0, "there has to be at least one repetition");
let mut rng = Rng::new(seed);
let planes = (0..shape.reps * shape.ksim * dim)
.map(|_| gauss(&mut rng))
.collect();
let scale = 1.0 / (shape.dproj as f32).sqrt();
let proj = (0..shape.reps * shape.dproj * dim)
.map(|_| {
if rng.next_u64() & 1 == 0 {
scale
} else {
-scale
}
})
.collect();
Encoder {
dim,
shape,
planes,
proj,
}
}
#[must_use]
pub fn dim(&self) -> usize {
self.dim
}
#[must_use]
pub fn shape(&self) -> Shape {
self.shape
}
#[must_use]
pub fn fde_dim(&self) -> usize {
self.shape.reps * self.buckets() * self.shape.dproj
}
#[must_use]
pub fn document(&self, tokens: &[f32]) -> Vec<f32> {
self.encode(tokens, true)
}
#[must_use]
pub fn query(&self, tokens: &[f32]) -> Vec<f32> {
self.encode(tokens, false)
}
fn buckets(&self) -> usize {
1usize << self.shape.ksim
}
fn encode(&self, tokens: &[f32], document: bool) -> Vec<f32> {
let dim = self.dim;
assert!(!tokens.is_empty(), "there is nothing to encode");
assert_eq!(
tokens.len() % dim,
0,
"the tokens are {} numbers, which is not a whole number of {dim} \
dimensional vectors",
tokens.len()
);
let n = tokens.len() / dim;
let buckets = self.buckets();
let mut out = vec![0.0f32; self.fde_dim()];
let mut codes = vec![0u32; n];
let mut totals = vec![0.0f32; buckets * dim];
let mut counts = vec![0u32; buckets];
let mut fill = vec![0.0f32; dim];
for r in 0..self.shape.reps {
totals.fill(0.0);
counts.fill(0);
for (t, code) in codes.iter_mut().enumerate() {
let x = &tokens[t * dim..(t + 1) * dim];
let k = self.bucket(r, x);
*code = k as u32;
counts[k] += 1;
for (into, c) in totals[k * dim..(k + 1) * dim].iter_mut().zip(x) {
*into += c;
}
}
for k in 0..buckets {
let at = (r * buckets + k) * self.shape.dproj;
if counts[k] > 0 {
let scale = if document {
1.0 / counts[k] as f32
} else {
1.0
};
self.project(r, &totals[k * dim..(k + 1) * dim], scale, at, &mut out);
} else if document {
let hits = nearest_by_hamming(tokens, dim, &codes, k as u32, &mut fill);
self.project(r, &fill, 1.0 / hits as f32, at, &mut out);
}
}
}
unit(&mut out);
out
}
fn bucket(&self, rep: usize, x: &[f32]) -> usize {
let planes = &self.planes[rep * self.shape.ksim * self.dim..];
let mut code = 0usize;
for b in 0..self.shape.ksim {
let plane = &planes[b * self.dim..(b + 1) * self.dim];
code |= usize::from(dot(plane, x) > 0.0) << b;
}
code
}
fn project(&self, rep: usize, x: &[f32], scale: f32, at: usize, out: &mut [f32]) {
let m = &self.proj[rep * self.shape.dproj * self.dim..];
for j in 0..self.shape.dproj {
out[at + j] = dot(&m[j * self.dim..(j + 1) * self.dim], x) * scale;
}
}
}
fn nearest_by_hamming(
tokens: &[f32],
dim: usize,
codes: &[u32],
want: u32,
fill: &mut [f32],
) -> u32 {
let mut best = u32::MAX;
let mut hits = 0u32;
for (t, code) in codes.iter().enumerate() {
let apart = (code ^ want).count_ones();
if apart > best {
continue;
}
if apart < best {
best = apart;
hits = 0;
fill.fill(0.0);
}
hits += 1;
for (into, c) in fill.iter_mut().zip(&tokens[t * dim..(t + 1) * dim]) {
*into += c;
}
}
hits
}
#[must_use]
pub fn chamfer(query: &[f32], doc: &[f32], dim: usize) -> f32 {
assert!(dim > 0, "a token has to have a dimension");
assert_eq!(query.len() % dim, 0, "the query is not whole tokens");
assert_eq!(doc.len() % dim, 0, "the document is not whole tokens");
if query.is_empty() || doc.is_empty() {
return 0.0;
}
let mut total = 0.0f32;
for q in query.chunks_exact(dim) {
let mut best = f32::NEG_INFINITY;
for p in doc.chunks_exact(dim) {
let d = dot(q, p);
if d > best {
best = d;
}
}
total += best;
}
total / (query.len() / dim) as f32
}
fn dot(a: &[f32], b: &[f32]) -> f32 {
let mut totals = [0.0f32; 8];
let mut i = 0;
while i + 8 <= a.len() {
for (k, total) in totals.iter_mut().enumerate() {
*total += a[i + k] * b[i + k];
}
i += 8;
}
let mut sum = 0.0f32;
for total in totals {
sum += total;
}
while i < a.len() {
sum += a[i] * b[i];
i += 1;
}
sum
}
fn gauss(rng: &mut Rng) -> f32 {
let u1 = (uniform(rng)).max(f32::MIN_POSITIVE);
let u2 = uniform(rng);
(-2.0 * u1.ln()).sqrt() * (core::f32::consts::TAU * u2).cos()
}
fn uniform(rng: &mut Rng) -> f32 {
(rng.next_u64() >> 40) as f32 / (1u32 << 24) as f32
}
fn unit(v: &mut [f32]) {
let len = v.iter().map(|c| c * c).sum::<f32>().sqrt();
if len > 0.0 {
for c in v {
*c /= len;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{Bits, Partitions, Tuning, Vectors};
fn corpus(dim: usize, docs: usize, tokens: usize, topics: usize, seed: u64) -> Vec<Vec<f32>> {
let mut rng = Rng::new(seed);
let concepts: Vec<Vec<f32>> = (0..topics).map(|_| draw(dim, &mut rng)).collect();
(0..docs)
.map(|_| {
let about: Vec<usize> = (0..4).map(|_| rng.below(topics)).collect();
let mut out = Vec::with_capacity(tokens * dim);
for t in 0..tokens {
let base = &concepts[about[t % about.len()]];
out.extend_from_slice(&near(base, 0.5, &mut rng));
}
out
})
.collect()
}
fn queries(
docs: &[Vec<f32>],
dim: usize,
n: usize,
len: usize,
seed: u64,
) -> Vec<(usize, Vec<f32>)> {
let mut rng = Rng::new(seed);
(0..n)
.map(|_| {
let from = rng.below(docs.len());
let doc = &docs[from];
let have = doc.len() / dim;
let mut q = Vec::with_capacity(len * dim);
for _ in 0..len {
let t = rng.below(have);
let token = &doc[t * dim..(t + 1) * dim];
q.extend_from_slice(&near(token, 0.35, &mut rng));
}
(from, q)
})
.collect()
}
fn near(base: &[f32], off: f32, rng: &mut Rng) -> Vec<f32> {
let noise = draw(base.len(), rng);
let mut v: Vec<f32> = base.iter().zip(&noise).map(|(c, o)| c + o * off).collect();
unit(&mut v);
v
}
fn draw(dim: usize, rng: &mut Rng) -> Vec<f32> {
let mut v: Vec<f32> = (0..dim).map(|_| gauss(rng)).collect();
unit(&mut v);
v
}
fn dot(a: &[f32], b: &[f32]) -> f32 {
a.iter().zip(b).map(|(x, y)| x * y).sum()
}
fn found(dim: usize, shape: Shape, keep: usize, seed: u64) -> f32 {
let enc = Encoder::new(dim, shape, seed);
let docs = corpus(dim, 300, 24, 64, seed);
let qs = queries(&docs, dim, 40, 6, seed ^ 0x5eed);
let fdes: Vec<Vec<f32>> = docs.iter().map(|d| enc.document(d)).collect();
let mut hits = 0usize;
for (from, q) in &qs {
let f = enc.query(q);
let mut by: Vec<(usize, f32)> = fdes
.iter()
.enumerate()
.map(|(i, d)| (i, dot(&f, d)))
.collect();
by.select_nth_unstable_by(keep, |a, b| b.1.total_cmp(&a.1));
hits += usize::from(by[..keep].iter().any(|(i, _)| i == from));
}
hits as f32 / qs.len() as f32
}
#[test]
fn an_encoding_is_the_length_it_says_it_is() {
let shape = Shape {
ksim: 3,
dproj: 8,
reps: 4,
};
let enc = Encoder::new(32, shape, 1);
assert_eq!(enc.fde_dim(), 4 * 8 * 8);
assert_eq!(enc.dim(), 32);
for tokens in [1usize, 2, 40] {
let set = corpus(32, 1, tokens, 4, 9).remove(0);
assert_eq!(enc.document(&set).len(), enc.fde_dim());
assert_eq!(enc.query(&set).len(), enc.fde_dim());
}
}
#[test]
fn the_same_seed_is_the_same_encoder() {
let set = corpus(24, 1, 9, 4, 3).remove(0);
let a = Encoder::new(24, Shape::default(), 77).document(&set);
let b = Encoder::new(24, Shape::default(), 77).document(&set);
let other = Encoder::new(24, Shape::default(), 78).document(&set);
assert_eq!(a, b);
assert_ne!(a, other);
}
#[test]
fn chamfer_is_the_best_match_for_each_query_token() {
let dim = 4;
let doc = [1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0];
let half = core::f32::consts::FRAC_1_SQRT_2;
let query = [0.0, 1.0, 0.0, 0.0, half, 0.0, half, 0.0];
assert!((chamfer(&query, &doc, dim) - (1.0 + half) / 2.0).abs() < 1e-6);
let one = [1.0f32, 0.0, 0.0, 0.0];
assert!((chamfer(&one, &doc, dim) - 1.0).abs() < 1e-6);
assert!((chamfer(&doc, &one, dim) - 0.5).abs() < 1e-6);
assert_eq!(chamfer(&[], &doc, dim), 0.0);
}
#[test]
fn a_one_token_document_fills_every_bucket() {
let dim = 16;
let shape = Shape {
ksim: 3,
dproj: 8,
reps: 2,
};
let enc = Encoder::new(dim, shape, 11);
let one = corpus(dim, 1, 1, 4, 5).remove(0);
let fde = enc.document(&one);
for r in 0..shape.reps {
let first = &fde[r * 8 * shape.dproj..][..shape.dproj];
assert!(first.iter().any(|c| c.abs() > 1e-6), "rep {r} is empty");
for k in 1..8 {
let block = &fde[(r * 8 + k) * shape.dproj..][..shape.dproj];
for (a, b) in first.iter().zip(block) {
assert!((a - b).abs() < 1e-6, "rep {r} bucket {k} differs");
}
}
}
}
#[test]
fn the_encoding_finds_the_document_a_query_came_from() {
let got = found(48, Shape::default(), 10, 4242);
assert!(
got >= 0.9,
"the encoding's top ten held it {got} of the time"
);
}
#[test]
fn more_repetitions_find_it_more_often() {
let one = found(
48,
Shape {
reps: 1,
..Shape::default()
},
1,
4242,
);
let many = found(
48,
Shape {
reps: 16,
..Shape::default()
},
1,
4242,
);
assert!(
many > one + 0.1,
"sixteen repetitions found it {many} of the time against one repetition's {one}"
);
}
struct Fdes(Vec<Vec<f32>>);
impl Vectors for Fdes {
fn get(&self, id: u64, into: &mut [f32]) -> bool {
match self.0.get(id as usize) {
Some(v) => {
into.copy_from_slice(v);
true
}
None => false,
}
}
}
#[test]
fn retrieval_then_a_chamfer_rerank_finds_the_right_document() {
let dim = 48;
let enc = Encoder::new(dim, Shape::default(), 909);
let docs = corpus(dim, 300, 24, 64, 909);
let qs = queries(&docs, dim, 40, 6, 0xbeef);
let fdes = Fdes(docs.iter().map(|d| enc.document(d)).collect());
let mut ix = Partitions::new(
enc.fde_dim(),
Bits::One,
7,
Tuning {
posting: 48,
..Tuning::default()
},
);
for (id, f) in fdes.0.iter().enumerate() {
ix.insert(id as u64, f);
}
ix.maintain(&fdes, 1 << 20);
let mut first = 0usize;
for (from, q) in &qs {
let best = ix
.search(&enc.query(q), 40, &fdes)
.into_iter()
.map(|h| (h.id as usize, chamfer(q, &docs[h.id as usize], dim)))
.max_by(|a, b| a.1.total_cmp(&b.1));
first += usize::from(best.map(|(id, _)| id) == Some(*from));
}
let got = first as f32 / qs.len() as f32;
assert!(
got >= 0.9,
"the right document came first {got} of the time"
);
}
#[test]
fn the_rerank_is_the_only_chamfer_anyone_pays_for() {
let dim = 32;
let enc = Encoder::new(dim, Shape::default(), 5);
let docs = corpus(dim, 200, 24, 32, 5);
let (from, q) = queries(&docs, dim, 1, 6, 17).remove(0);
let scanned: usize = docs.iter().map(|d| d.len() / dim).sum::<usize>() * (q.len() / dim);
let fdes: Vec<Vec<f32>> = docs.iter().map(|d| enc.document(d)).collect();
let f = enc.query(&q);
let mut by: Vec<(usize, f32)> = fdes
.iter()
.enumerate()
.map(|(i, d)| (i, dot(&f, d)))
.collect();
by.sort_by(|a, b| b.1.total_cmp(&a.1));
let kept: Vec<usize> = by[..20].iter().map(|(i, _)| *i).collect();
let reranked: usize =
kept.iter().map(|i| docs[*i].len() / dim).sum::<usize>() * (q.len() / dim);
assert!(
kept.contains(&from),
"the document it came from was dropped"
);
assert!(
reranked * 8 < scanned,
"reranking {reranked} token pairs against a full scan's {scanned} is not a saving"
);
}
}