#![cfg(test)]
use crate::dist::sqdist;
use crate::rabitq::{Bits, Coded, Quantizer};
use yo_common::Rng;
fn clumped(n: usize, dim: usize, seed: u64) -> Vec<f32> {
let mut rng = Rng::new(seed);
let mut unit = || (rng.next_u64() >> 40) as f32 / (1u32 << 24) as f32;
let groups = 24;
let hubs: Vec<f32> = (0..groups * dim).map(|_| unit()).collect();
let mut out = Vec::with_capacity(n * dim);
for i in 0..n {
let h = (i % groups) * dim;
for d in 0..dim {
out.push(hubs[h + d] + (unit() - 0.5) * 0.2);
}
}
out
}
fn truth(u: &[f32], centroids: &[f32], dim: usize, n: usize, want: usize) -> Vec<usize> {
let mut by: Vec<(usize, f32)> = (0..n)
.map(|p| (p, sqdist(u, ¢roids[p * dim..(p + 1) * dim])))
.collect();
by.sort_by(|a, b| a.1.total_cmp(&b.1));
by.truncate(want);
by.into_iter().map(|(p, _)| p).collect()
}
struct Codes {
quant: Quantizer,
mean: Vec<f32>,
codes: Vec<u8>,
meta: Vec<Coded>,
}
impl Codes {
fn build(centroids: &[f32], dim: usize, n: usize, bits: Bits) -> Codes {
let quant = Quantizer::new(dim, bits, 0x51de_0001);
let mut mean = vec![0.0f32; dim];
for p in 0..n {
for (m, c) in mean.iter_mut().zip(¢roids[p * dim..(p + 1) * dim]) {
*m += *c;
}
}
for m in &mut mean {
*m /= n as f32;
}
let width = quant.code_bytes();
let mut codes = vec![0u8; n * width];
let meta: Vec<Coded> = (0..n)
.map(|p| {
quant.encode_rotated(
¢roids[p * dim..(p + 1) * dim],
&mean,
&mut codes[p * width..(p + 1) * width],
)
})
.collect();
Codes {
quant,
mean,
codes,
meta,
}
}
fn estimate(&self, u: &[f32], scores: &mut Vec<f32>) {
scores.clear();
scores.resize(self.meta.len(), 0.0);
self.quant
.query_rotated(u, &self.mean)
.scan(&self.codes, &self.meta, scores);
}
fn head(&self, u: &[f32], centroids: &[f32], want: usize, scores: &mut Vec<f32>) -> Vec<usize> {
let dim = self.quant.dim();
let n = self.meta.len();
self.estimate(u, scores);
let order = |a: &(usize, f32), b: &(usize, f32)| a.1.total_cmp(&b.1);
let wide = (want * 4).clamp(128, 1024).min(n);
let mut by: Vec<(usize, f32)> = scores.iter().copied().enumerate().collect();
by.select_nth_unstable_by(wide - 1, order);
by.truncate(wide);
for entry in &mut by {
entry.1 = sqdist(u, ¢roids[entry.0 * dim..(entry.0 + 1) * dim]);
}
by.select_nth_unstable_by(want - 1, order);
by.truncate(want);
by.sort_unstable_by(order);
by.into_iter().map(|(p, _)| p).collect()
}
}
fn full_head(u: &[f32], centroids: &[f32], dim: usize, n: usize, want: usize) -> Vec<usize> {
let order = |a: &(usize, f32), b: &(usize, f32)| a.1.total_cmp(&b.1);
let mut by: Vec<(usize, f32)> = (0..n)
.map(|p| (p, sqdist(u, ¢roids[p * dim..(p + 1) * dim])))
.collect();
by.select_nth_unstable_by(want - 1, order);
by.truncate(want);
by.sort_unstable_by(order);
by.into_iter().map(|(p, _)| p).collect()
}
#[test]
#[cfg_attr(
miri,
ignore = "the count is the claim: an exact shortlist of sixteen out of two thousand centroids, a hundred queries deep and at two widths, and none of that survives being made small"
)]
fn coding_the_centroids_would_have_given_the_right_answer() {
for dim in [32usize, 128] {
let n = 2000;
let centroids = clumped(n, dim, 2);
let coded = Codes::build(¢roids, dim, n, Bits::Four);
let queries = clumped(100, dim, 3);
let mut scores = Vec::new();
for i in 0..100 {
let u = &queries[i * dim..(i + 1) * dim];
assert_eq!(
coded.head(u, ¢roids, 16, &mut scores),
truth(u, ¢roids, dim, n, 16),
"at dimension {dim}, query {i}"
);
}
}
}
#[test]
#[ignore = "prints a table rather than asserting anything"]
fn what_ranking_the_centroids_costs() {
println!("centroids dimensions read in full coded and reranked");
for &(n, dim) in &[(2963usize, 1024usize), (2963, 128), (10_000, 768)] {
let centroids = clumped(n, dim, 2);
let coded = Codes::build(¢roids, dim, n, Bits::Four);
let queries = clumped(200, dim, 3);
let want = 128;
let mut scores = Vec::new();
let mut sink = 0usize;
for i in 0..200 {
let u = &queries[i * dim..(i + 1) * dim];
sink += full_head(u, ¢roids, dim, n, want).len();
sink += coded.head(u, ¢roids, want, &mut scores).len();
}
let at = std::time::Instant::now();
for i in 0..200 {
sink += full_head(&queries[i * dim..(i + 1) * dim], ¢roids, dim, n, want).len();
}
let full = at.elapsed().as_nanos() as f64 / 200.0 / 1000.0;
let at = std::time::Instant::now();
for i in 0..200 {
sink += coded
.head(
&queries[i * dim..(i + 1) * dim],
¢roids,
want,
&mut scores,
)
.len();
}
let bits = at.elapsed().as_nanos() as f64 / 200.0 / 1000.0;
println!("{n:9} {dim:11} {full:11.1} us {bits:16.1} us ({sink})");
}
}
#[test]
#[ignore = "prints a table rather than asserting anything"]
fn how_much_slack_the_shortlist_needs() {
let n = 3000;
println!("dimensions head of one bit four bit");
for dim in [128usize, 1024] {
let centroids = clumped(n, dim, 2);
let queries = clumped(100, dim, 3);
for want in [16usize, 128, 512] {
let mut reach = Vec::new();
for bits in [Bits::One, Bits::Four] {
let coded = Codes::build(¢roids, dim, n, bits);
let mut worst = 0usize;
let mut scores = Vec::new();
for i in 0..100 {
let u = &queries[i * dim..(i + 1) * dim];
let head = truth(u, ¢roids, dim, n, want);
coded.estimate(u, &mut scores);
let mut order: Vec<usize> = (0..n).collect();
order.sort_by(|&a, &b| scores[a].total_cmp(&scores[b]));
let mut place = vec![0usize; n];
for (rank, &p) in order.iter().enumerate() {
place[p] = rank;
}
for p in &head {
worst = worst.max(place[*p] + 1);
}
}
reach.push(worst);
}
println!("{dim:10} {want:8} {:9} {:10}", reach[0], reach[1]);
}
}
}