use crate::sift::DESC_LEN;
use rayon::prelude::*;
pub struct VocabParams {
pub branching: usize,
pub depth: usize,
pub sample: usize,
pub iters: usize,
pub max_paths: usize,
pub path_ratio: f32,
pub seed: u64,
}
const DESC_PER_WORD: usize = 3;
impl VocabParams {
pub fn for_corpus(n_desc: usize) -> VocabParams {
let p = VocabParams::default();
let target = (n_desc / DESC_PER_WORD).max(2);
let mut depth = 1usize;
while depth < 6 && p.branching.pow(depth as u32) < target {
depth += 1;
}
let mut branching = 2usize;
while branching < p.branching && branching.pow(depth as u32) < target {
branching += 1;
}
VocabParams { depth, branching, ..p }
}
}
impl Default for VocabParams {
fn default() -> Self {
VocabParams {
branching: 16,
depth: 4,
sample: 160_000,
iters: 8,
max_paths: 3,
path_ratio: 1.3,
seed: 0x5eed_1234_abcd_ef01,
}
}
}
pub struct Vocabulary {
pub branching: usize,
pub depth: usize,
levels: Vec<Vec<u8>>,
head: Vec<Vec<Kids>>,
node_of: Vec<Vec<u32>>,
max_paths: usize,
path_ratio: f32,
}
#[derive(Clone, Copy, Default)]
struct Kids {
first: u32,
n: u32,
}
const DEAD: u32 = u32::MAX;
const FRONTIER: usize = 64;
const SPAN: usize = 32;
const LANES: usize = 8;
const MAX_BRANCH: usize = 16;
#[inline]
fn d2(a: &[f32], b: &[f32]) -> f32 {
let mut s = 0.0;
for i in 0..DESC_LEN {
let d = a[i] - b[i];
s += d * d;
}
s
}
struct Rng(u64);
impl Rng {
fn next(&mut self) -> u64 {
let mut x = self.0;
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
self.0 = x;
x
}
fn below(&mut self, n: usize) -> usize {
(self.next() % n.max(1) as u64) as usize
}
}
#[inline]
fn prefetch_centres(centres: &[u8], kids: Kids) {
if kids.n == 0 {
return;
}
#[cfg(target_arch = "x86_64")]
unsafe {
use std::arch::x86_64::{_mm_prefetch, _MM_HINT_T0};
let start = kids.first as usize * DESC_LEN;
let end = (start + kids.n as usize * DESC_LEN).min(centres.len());
let mut at = start;
let mut lines = 0;
while at < end && lines < 4 {
_mm_prefetch(centres.as_ptr().add(at) as *const i8, _MM_HINT_T0);
at += 64;
lines += 1;
}
}
#[cfg(not(target_arch = "x86_64"))]
let _ = centres;
}
impl Vocabulary {
pub fn n_words(&self) -> usize {
self.branching.pow(self.depth as u32)
}
pub fn n_live_words(&self) -> usize {
self.node_of[self.depth - 1].len()
}
pub fn build(descriptors: &[u8], p: &VocabParams) -> Vocabulary {
assert!(p.max_paths * p.branching <= FRONTIER, "vocabulary frontier too small");
assert!(p.branching <= MAX_BRANCH, "branching wider than k-means' accumulators");
let n = descriptors.len() / DESC_LEN;
let mut rng = Rng(p.seed);
let take = p.sample.min(n);
let mut idx: Vec<u32> = (0..n as u32).collect();
for i in 0..take {
let j = i + rng.below(n - i);
idx.swap(i, j);
}
let sample: Vec<f32> = idx[..take]
.iter()
.flat_map(|&i| {
descriptors[i as usize * DESC_LEN..(i as usize + 1) * DESC_LEN]
.iter()
.map(|&v| v as f32)
})
.collect();
let mut levels: Vec<Vec<u8>> = Vec::with_capacity(p.depth);
let mut heads: Vec<Vec<Kids>> = Vec::with_capacity(p.depth);
let mut node_ofs: Vec<Vec<u32>> = Vec::with_capacity(p.depth);
let mut assign: Vec<u32> = vec![0; take];
let mut parents = 1usize;
for _level in 0..p.depth {
let nodes = parents * p.branching;
let mut centres: Vec<u8> = Vec::new();
let mut slot = vec![DEAD; nodes];
let mut groups: Vec<Vec<u32>> = vec![Vec::new(); parents];
for (i, &a) in assign.iter().enumerate() {
groups[a as usize].push(i as u32);
}
let results: Vec<(usize, Vec<f32>, Vec<bool>, Vec<(u32, u32)>)> = groups
.par_iter()
.enumerate()
.map(|(g, members)| {
let (c, l, a) = kmeans(&sample, members, p.branching, p.iters, p.seed ^ (g as u64 + 1));
(g, c, l, a)
})
.collect();
for (g, c, l, a) in results {
let base = g * p.branching;
let live: Vec<usize> = (0..p.branching).filter(|&c| l[c]).collect();
let first = (centres.len() / DESC_LEN) as u32;
for (k, &child) in live.iter().enumerate() {
slot[base + child] = first + k as u32;
}
centres.reserve(live.len() * DESC_LEN);
for k in 0..live.len() {
for i in 0..DESC_LEN {
centres.push(quantise_centre(c[k * DESC_LEN + i]));
}
}
for (i, child) in a {
assign[i as usize] = (base + child as usize) as u32;
}
}
centres.shrink_to_fit();
let mut head = vec![Kids::default(); parents];
let mut node_of = vec![0u32; centres.len() / DESC_LEN];
for parent in 0..parents {
let base = parent * p.branching;
let mut first = u32::MAX;
let mut n = 0u32;
for c in 0..p.branching {
let sl = slot[base + c];
if sl == DEAD {
continue;
}
if first == u32::MAX {
first = sl;
}
node_of[sl as usize] = (base + c) as u32;
n += 1;
}
head[parent] = Kids { first: if first == u32::MAX { 0 } else { first }, n };
}
levels.push(centres);
heads.push(head);
node_ofs.push(node_of);
parents = nodes;
}
Vocabulary {
branching: p.branching,
depth: p.depth,
levels,
head: heads,
node_of: node_ofs,
max_paths: p.max_paths,
path_ratio: p.path_ratio,
}
}
pub fn quantise(&self, desc: &[u8], out: &mut Vec<u32>) {
out.clear();
let q: &[u8; DESC_LEN] = desc[..DESC_LEN].try_into().unwrap();
let mut cur = [(0u32, 0u32); FRONTIER];
let mut next = [(0u32, 0u32); FRONTIER];
let mut n_cur = 1usize;
for l in 0..self.depth {
let mut n_next = 0usize;
let centres = &self.levels[l];
let head = &self.head[l];
let mut par = [0u32; FRONTIER];
if l == 0 {
par[0] = 0;
} else {
let above = &self.node_of[l - 1];
for i in 0..n_cur {
par[i] = above[cur[i].0 as usize];
}
}
let mut kids = [Kids::default(); FRONTIER];
for pi in 0..n_cur {
kids[pi] = head[par[pi] as usize];
prefetch_centres(centres, kids[pi]);
}
let keep = self.max_paths;
let mut held = 0usize;
for pi in 0..n_cur {
let kids = kids[pi];
let n_live = kids.n as usize;
if n_live == 0 {
continue;
}
let first = kids.first as usize;
let blk = ¢res[first * DESC_LEN..(first + n_live) * DESC_LEN];
for (k, c) in blk.chunks_exact(DESC_LEN).enumerate() {
if n_next == FRONTIER {
break;
}
let e = ((first + k) as u32, dist2(q, c));
next[n_next] = e;
n_next += 1;
if held == keep && !(e.1 < cur[held - 1].1) {
continue;
}
let mut j = held.min(keep - 1);
while j > 0 && e.1 < cur[j - 1].1 {
cur[j] = cur[j - 1];
j -= 1;
}
cur[j] = e;
held += (held < keep) as usize;
}
}
if n_next == 0 {
break;
}
let mut ambiguous = false;
for j in 1..held {
ambiguous |= cur[j].1 == cur[j - 1].1;
}
if held < n_next {
let bound = cur[held - 1].1;
let mut n_eq = 0usize;
for i in 0..n_next {
n_eq += (next[i].1 == bound) as usize;
}
ambiguous |= n_eq > 1;
}
if ambiguous {
let nx = &mut next[..n_next];
nx.sort_unstable_by_key(|e| e.1);
cur[..held].copy_from_slice(&next[..held]);
}
let cut = cur[0].1 as f32 * self.path_ratio * self.path_ratio + 1.0;
n_cur = held;
while n_cur > 1 && (cur[n_cur - 1].1 as f32) > cut {
n_cur -= 1;
}
}
for &(slot, _) in cur[..n_cur].iter() {
out.push(slot);
}
}
}
#[inline]
fn quantise_centre(v: f32) -> u8 {
v.round() as u8
}
#[cfg(target_feature = "avx2")]
#[inline]
pub fn dist2(q: &[u8; DESC_LEN], c: &[u8]) -> u32 {
debug_assert!(c.len() >= DESC_LEN && DESC_LEN % 32 == 0);
unsafe {
use std::arch::x86_64::*;
let zero = _mm256_setzero_si256();
let mut acc = zero;
for o in (0..DESC_LEN).step_by(32) {
let a = _mm256_loadu_si256(q.as_ptr().add(o) as *const __m256i);
let b = _mm256_loadu_si256(c.as_ptr().add(o) as *const __m256i);
let d = _mm256_or_si256(_mm256_subs_epu8(a, b), _mm256_subs_epu8(b, a));
let lo = _mm256_unpacklo_epi8(d, zero);
let hi = _mm256_unpackhi_epi8(d, zero);
acc = _mm256_add_epi32(acc, _mm256_madd_epi16(lo, lo));
acc = _mm256_add_epi32(acc, _mm256_madd_epi16(hi, hi));
}
let half = _mm_add_epi32(_mm256_castsi256_si128(acc), _mm256_extracti128_si256(acc, 1));
let pair = _mm_add_epi32(half, _mm_shuffle_epi32(half, 0b00_00_11_10));
let one = _mm_add_epi32(pair, _mm_shuffle_epi32(pair, 0b00_00_00_01));
_mm_cvtsi128_si32(one) as u32
}
}
#[cfg(not(target_feature = "avx2"))]
#[inline]
pub fn dist2(q: &[u8; DESC_LEN], c: &[u8]) -> u32 {
const LANES: usize = 16;
let mut acc = [0u32; LANES];
for (a, b) in q.chunks_exact(LANES).zip(c[..DESC_LEN].chunks_exact(LANES)) {
for l in 0..LANES {
let d = a[l] as i32 - b[l] as i32;
acc[l] += (d * d) as u32;
}
}
let mut s = 0u32;
for l in 0..LANES {
s += acc[l];
}
s
}
fn child_dists(k: usize, q: &[f32], blk: &[f32], acc: &mut [f32; MAX_BRANCH]) {
macro_rules! widths {
($($w:literal)*) => {
match k {
$($w => {
let mut a = [0f32; $w];
if $w % LANES == 0 {
dists::<$w>(q, blk, &mut a);
} else {
for c in 0..DESC_LEN / SPAN {
let d0 = c * SPAN;
dists::<$w>(&q[d0..d0 + SPAN], &blk[d0 * $w..(d0 + SPAN) * $w], &mut a);
}
}
acc[..$w].copy_from_slice(&a);
})*
_ => dists_dyn(k, q, blk, acc),
}
};
}
widths!(1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16);
}
#[inline(always)]
fn dists<const K: usize>(q: &[f32], blk: &[f32], acc: &mut [f32; K]) {
for (&qi, row) in q.iter().zip(blk.chunks_exact(K)) {
for (a, &c) in acc.iter_mut().zip(row.iter()) {
let d = qi - c;
*a += d * d;
}
}
}
fn dists_dyn(k: usize, q: &[f32], blk: &[f32], acc: &mut [f32; MAX_BRANCH]) {
for (&qi, row) in q.iter().zip(blk.chunks_exact(k)) {
for (a, &c) in acc[..k].iter_mut().zip(row.iter()) {
let d = qi - c;
*a += d * d;
}
}
}
fn kmeans(data: &[f32], members: &[u32], k: usize, iters: usize, seed: u64) -> (Vec<f32>, Vec<bool>, Vec<(u32, u32)>) {
assert!(k <= MAX_BRANCH);
let mut live = vec![false; k];
let mut assign: Vec<(u32, u32)> = Vec::with_capacity(members.len());
if members.is_empty() {
return (Vec::new(), live, assign);
}
if members.len() <= k {
let mut centres = Vec::with_capacity(members.len() * DESC_LEN);
for (c, &m) in members.iter().enumerate() {
centres.extend_from_slice(&data[m as usize * DESC_LEN..(m as usize + 1) * DESC_LEN]);
live[c] = true;
assign.push((m, c as u32));
}
return (centres, live, assign);
}
let mut centres = vec![0f32; k * DESC_LEN];
let mut rng = Rng(seed | 1);
let mut chosen: Vec<u32> = Vec::with_capacity(k);
chosen.push(members[rng.below(members.len())]);
let mut best_d: Vec<f32> = members
.par_iter()
.map(|&m| {
d2(
&data[m as usize * DESC_LEN..(m as usize + 1) * DESC_LEN],
&data[chosen[0] as usize * DESC_LEN..(chosen[0] as usize + 1) * DESC_LEN],
)
})
.collect();
while chosen.len() < k {
let total: f64 = best_d.iter().map(|&v| v as f64).sum();
if total <= 0.0 {
break;
}
let mut t = (rng.next() as f64 / u64::MAX as f64) * total;
let mut pick = members.len() - 1;
for (i, &v) in best_d.iter().enumerate() {
t -= v as f64;
if t <= 0.0 {
pick = i;
break;
}
}
let c = members[pick];
chosen.push(c);
let cd = &data[c as usize * DESC_LEN..(c as usize + 1) * DESC_LEN];
best_d
.par_iter_mut()
.zip(members.par_iter())
.for_each(|(b, &m)| {
let d = d2(&data[m as usize * DESC_LEN..(m as usize + 1) * DESC_LEN], cd);
if d < *b {
*b = d;
}
});
}
let kk = chosen.len();
for (c, &m) in chosen.iter().enumerate() {
centres[c * DESC_LEN..(c + 1) * DESC_LEN]
.copy_from_slice(&data[m as usize * DESC_LEN..(m as usize + 1) * DESC_LEN]);
}
let mut owner = vec![0u32; members.len()];
let mut tc = vec![0f32; kk * DESC_LEN];
for it in 0..iters {
for c in 0..kk {
for i in 0..DESC_LEN {
tc[i * kk + c] = centres[c * DESC_LEN + i];
}
}
let changed: usize = owner
.par_iter_mut()
.zip(members.par_iter())
.map(|(own, &m)| {
let dv = &data[m as usize * DESC_LEN..(m as usize + 1) * DESC_LEN];
let mut acc = [0f32; MAX_BRANCH];
child_dists(kk, dv, &tc, &mut acc);
let mut best = (f32::MAX, 0u32);
for c in 0..kk {
if acc[c] < best.0 {
best = (acc[c], c as u32);
}
}
let same = *own == best.1;
*own = best.1;
!same as usize
})
.sum();
if changed == 0 && it > 0 {
break;
}
let mut sums = vec![0f64; kk * DESC_LEN];
let mut counts = vec![0u32; kk];
for (i, &m) in members.iter().enumerate() {
let c = owner[i] as usize;
counts[c] += 1;
let dv = &data[m as usize * DESC_LEN..(m as usize + 1) * DESC_LEN];
let s = &mut sums[c * DESC_LEN..(c + 1) * DESC_LEN];
for j in 0..DESC_LEN {
s[j] += dv[j] as f64;
}
}
for c in 0..kk {
if counts[c] == 0 {
continue;
}
let inv = 1.0 / counts[c] as f64;
for j in 0..DESC_LEN {
centres[c * DESC_LEN + j] = (sums[c * DESC_LEN + j] * inv) as f32;
}
}
}
let mut used = vec![0u32; kk];
for &c in owner.iter() {
used[c as usize] += 1;
}
for c in 0..kk {
live[c] = used[c] > 0;
}
for (i, &m) in members.iter().enumerate() {
assign.push((m, owner[i]));
}
let n_live = live.iter().filter(|&&l| l).count();
if n_live < k {
let mut compact = Vec::with_capacity(n_live * DESC_LEN);
for c in 0..kk {
if live[c] {
compact.extend_from_slice(¢res[c * DESC_LEN..(c + 1) * DESC_LEN]);
}
}
centres = compact;
}
(centres, live, assign)
}
#[derive(Clone, Debug, Default)]
pub struct WordList {
pub word: Vec<u32>,
pub kp: Vec<u32>,
}
impl WordList {
pub fn runs(&self) -> impl Iterator<Item = (u32, u32)> + '_ {
self.runs_at().map(|(_, w, c)| (w, c))
}
pub fn runs_at(&self) -> Runs<'_> {
Runs { wl: self, i: 0 }
}
}
pub struct Runs<'a> {
wl: &'a WordList,
i: usize,
}
impl Iterator for Runs<'_> {
type Item = (usize, u32, u32);
fn next(&mut self) -> Option<(usize, u32, u32)> {
if self.i >= self.wl.word.len() {
return None;
}
let at = self.i;
let w = self.wl.word[at];
let mut c = 0u32;
while self.i < self.wl.word.len() && self.wl.word[self.i] == w {
c += 1;
self.i += 1;
}
Some((at, w, c))
}
}
pub struct InvertedFile {
off: Vec<u32>,
data: Vec<(u32, u32)>,
idf: Vec<f32>,
}
const POST_AHEAD: usize = 6;
const POST_LINES: usize = 8;
#[inline]
fn prefetch_postings(inv: &InvertedFile, words: &[u32], at: usize) {
let Some(&w) = words.get(at) else { return };
#[cfg(target_arch = "x86_64")]
unsafe {
use std::arch::x86_64::{_mm_prefetch, _MM_HINT_T0};
let w = w as usize;
let start = *inv.off.get_unchecked(w) as usize;
let end = *inv.off.get_unchecked(w + 1) as usize;
if start >= end {
return;
}
let p = inv.data.as_ptr().add(start) as *const i8;
let bytes = (end - start) * std::mem::size_of::<(u32, u32)>();
let mut at = 0usize;
let mut lines = 0usize;
while at < bytes && lines < POST_LINES {
_mm_prefetch(p.add(at), _MM_HINT_T0);
at += 64;
lines += 1;
}
}
#[cfg(not(target_arch = "x86_64"))]
let _ = (inv, w);
}
impl InvertedFile {
pub fn build(lists: &[WordList], n_words: usize, max_posting: usize) -> InvertedFile {
let n = lists.len();
let mut df = vec![0u32; n_words];
for wl in lists.iter() {
for (w, _) in wl.runs() {
df[w as usize] += 1;
}
}
let mut idf = vec![0f32; n_words];
let mut off: Vec<u32> = Vec::with_capacity(n_words + 1);
let mut total = 0u32;
for w in 0..n_words {
off.push(total);
let d = df[w] as usize;
if d == 0 || d > max_posting {
continue;
}
idf[w] = (n as f32 / d as f32).ln();
total += df[w];
}
off.push(total);
let mut data = vec![(0u32, 0u32); total as usize];
let mut cursor = df;
cursor.copy_from_slice(&off[..n_words]);
for (img, wl) in lists.iter().enumerate() {
for (w, c) in wl.runs() {
let w = w as usize;
if off[w + 1] == off[w] {
continue;
}
let at = cursor[w] as usize;
data[at] = (img as u32, c);
cursor[w] = at as u32 + 1;
}
}
InvertedFile { off, data, idf }
}
pub fn query(
&self,
wl: &WordList,
exclude: u32,
acc: &mut [f32],
touched: &mut Vec<u32>,
out: &mut Vec<(u32, f32)>,
) {
out.clear();
touched.clear();
let mut qmass = 0f32;
let words = &wl.word[..];
for (at, w, c) in wl.runs_at() {
let w = w as usize;
prefetch_postings(self, words, at + POST_AHEAD);
let post = &self.data[self.off[w] as usize..self.off[w + 1] as usize];
if post.is_empty() {
continue;
}
let idf2 = self.idf[w] * self.idf[w];
qmass += c as f32 * idf2;
for &(other, cnt) in post.iter() {
if other == exclude {
continue;
}
let a = &mut acc[other as usize];
if *a == 0.0 {
touched.push(other);
}
*a += cnt.min(c) as f32 * idf2;
}
}
let inv = 1.0 / qmass.max(1e-6);
for &o in touched.iter() {
let s = acc[o as usize] * inv;
acc[o as usize] = 0.0;
out.push((o, s));
}
}
}
#[cfg(target_feature = "avx2")]
#[inline]
fn blocks_meet(a: &[u32], b: &[u32]) -> bool {
debug_assert!(a.len() >= BLOCK && b.len() >= BLOCK);
unsafe {
use std::arch::x86_64::*;
let bv = _mm256_loadu_si256(b.as_ptr() as *const __m256i);
let mut acc = _mm256_setzero_si256();
for k in 0..BLOCK {
let av = _mm256_set1_epi32(*a.get_unchecked(k) as i32);
acc = _mm256_or_si256(acc, _mm256_cmpeq_epi32(av, bv));
}
_mm256_movemask_epi8(acc) != 0
}
}
#[cfg(not(target_feature = "avx2"))]
#[inline]
fn blocks_meet(_a: &[u32], _b: &[u32]) -> bool {
true
}
#[cfg(target_feature = "avx2")]
const BLOCK: usize = 8;
#[cfg(not(target_feature = "avx2"))]
const BLOCK: usize = 0;
pub fn shared(a: &WordList, b: &WordList, out: &mut Vec<(u32, u32)>, cap: usize) {
out.clear();
let mut hi = 0u32;
let (aw, bw) = (&a.word[..], &b.word[..]);
let (na, nb) = (aw.len(), bw.len());
let (mut i, mut j) = (0usize, 0usize);
while i < na && j < nb {
let (mut ia, mut jb) = (i + BLOCK, j + BLOCK);
while BLOCK > 0 && ia <= na && jb <= nb && !blocks_meet(&aw[i..ia], &bw[j..jb]) {
let (am, bm) = (aw[ia - 1], bw[jb - 1]);
if am <= bm {
i = ia;
ia += BLOCK;
}
if bm <= am {
j = jb;
jb += BLOCK;
}
}
if i >= na || j >= nb {
break;
}
let (aend, bend) = if BLOCK == 0 { (na, nb) } else { (ia.min(na), jb.min(nb)) };
while i < aend && j < bend {
let (av, bv) = (aw[i], bw[j]);
i += (av < bv) as usize;
j += (bv < av) as usize;
if av != bv {
continue;
}
let (i0, j0) = (i, j);
while i < na && aw[i] == av {
i += 1;
}
while j < nb && bw[j] == av {
j += 1;
}
if (i - i0) * (j - j0) > 64 {
continue;
}
for x in i0..i {
for y in j0..j {
let e = (a.kp[x], b.kp[y]);
hi |= e.0 | e.1;
out.push(e);
}
}
if out.len() > cap {
sort_pairs(out, hi);
out.dedup();
return;
}
}
}
sort_pairs(out, hi);
out.dedup();
}
const RADIX_BITS: usize = 10;
const RADIX_BUCKETS: usize = 1 << RADIX_BITS;
const RADIX_MIN: usize = 192;
struct SortScratch {
tmp: Vec<(u32, u32)>,
cnt: Vec<u32>,
}
thread_local! {
static SORT_SCRATCH: std::cell::RefCell<SortScratch> =
const { std::cell::RefCell::new(SortScratch { tmp: Vec::new(), cnt: Vec::new() }) };
}
fn sort_pairs(out: &mut Vec<(u32, u32)>, hi: u32) {
if out.len() < RADIX_MIN || hi >= RADIX_BUCKETS as u32 {
out.sort_unstable();
return;
}
SORT_SCRATCH.with(|s| {
let s = &mut *s.borrow_mut();
let n = out.len();
if s.tmp.len() < n {
s.tmp.resize(n, (0, 0));
}
if s.cnt.len() != RADIX_BUCKETS {
s.cnt.resize(RADIX_BUCKETS, 0);
}
let cnt = &mut s.cnt[..RADIX_BUCKETS];
let tmp = &mut s.tmp[..n];
for pass in 0..2 {
cnt.fill(0);
let (src, dst): (&[(u32, u32)], &mut [(u32, u32)]) =
if pass == 0 { (&out[..n], tmp) } else { (tmp, &mut out[..n]) };
for e in src.iter() {
let k = if pass == 0 { e.1 } else { e.0 } as usize;
cnt[k] += 1;
}
let mut run = 0u32;
for c in cnt.iter_mut() {
let take = *c;
*c = run;
run += take;
}
for &e in src.iter() {
let k = if pass == 0 { e.1 } else { e.0 } as usize;
dst[cnt[k] as usize] = e;
cnt[k] += 1;
}
}
});
}
#[cfg(test)]
mod tests {
use super::*;
fn shared_reference(a: &WordList, b: &WordList, out: &mut Vec<(u32, u32)>, cap: usize) {
out.clear();
let (mut i, mut j) = (0usize, 0usize);
while i < a.word.len() && j < b.word.len() {
match a.word[i].cmp(&b.word[j]) {
std::cmp::Ordering::Less => i += 1,
std::cmp::Ordering::Greater => j += 1,
std::cmp::Ordering::Equal => {
let w = a.word[i];
let i0 = i;
while i < a.word.len() && a.word[i] == w {
i += 1;
}
let j0 = j;
while j < b.word.len() && b.word[j] == w {
j += 1;
}
if (i - i0) * (j - j0) > 64 {
continue;
}
for x in i0..i {
for y in j0..j {
out.push((a.kp[x], b.kp[y]));
}
}
if out.len() > cap {
break;
}
}
}
}
out.sort_unstable();
out.dedup();
}
#[test]
fn the_block_filter_intersects_exactly_as_a_merge_does() {
let mut rng = Lcg(0x243f_6a88_85a3_08d3);
let mut made = |len: usize, span: u32, reps: u32| -> WordList {
let mut pairs: Vec<(u32, u32)> = (0..len)
.map(|k| {
let w = (rng.byte() as u32) << 8 | rng.byte() as u32;
(w % span.max(1) / reps.max(1) * reps.max(1), k as u32)
})
.collect();
pairs.sort_unstable();
WordList {
word: pairs.iter().map(|p| p.0).collect(),
kp: pairs.iter().map(|p| p.1).collect(),
}
};
let mut got = Vec::new();
let mut want = Vec::new();
for &(la, lb) in [(0usize, 7usize), (1, 1), (3, 40), (9, 9), (17, 8), (64, 64), (300, 290), (1000, 30), (1500, 1200)].iter() {
for &(span, reps) in [(65535u32, 1u32), (400, 1), (64, 1), (65535, 7), (200, 3)].iter() {
let a = made(la, span, reps);
let b = made(lb, span, reps);
for cap in [60_000usize, 32, 3] {
shared(&a, &b, &mut got, cap);
shared_reference(&a, &b, &mut want, cap);
assert_eq!(got, want, "la {la} lb {lb} span {span} reps {reps} cap {cap}");
}
}
}
}
struct Lcg(u64);
impl Lcg {
fn byte(&mut self) -> u8 {
self.0 = self.0.wrapping_mul(6364136223846793005).wrapping_add(1442695040888963407);
(self.0 >> 33) as u8
}
}
#[test]
fn a_two_level_vocabulary_separates_what_is_separable() {
const GROUPS: usize = 16;
const SUBS: usize = 16;
const EACH: usize = 10;
let mut rng = Lcg(0x9e37_79b9_7f4a_7c15);
let mut desc: Vec<u8> = Vec::with_capacity(GROUPS * SUBS * EACH * DESC_LEN);
for _ in 0..GROUPS {
let centre: Vec<i32> = (0..DESC_LEN).map(|_| (rng.byte() / 2) as i32).collect();
for sub in 0..SUBS {
let mut mode = centre.clone();
for d in 0..8 {
mode[(sub * 8 + d) % DESC_LEN] += 40;
}
for _ in 0..EACH {
for &m in mode.iter() {
desc.push((m + (rng.byte() % 3) as i32 - 1).clamp(0, 255) as u8);
}
}
}
}
let p = VocabParams { depth: 2, ..Default::default() };
let v = Vocabulary::build(&desc, &p);
let mut words: Vec<Vec<u32>> = Vec::new();
let mut out = Vec::new();
for (i, d) in desc.chunks_exact(DESC_LEN).enumerate() {
v.quantise(d, &mut out);
assert!(!out.is_empty(), "descriptor {i} quantised to nothing");
assert!(out.iter().all(|&w| (w as usize) < v.n_words()));
words.push(out.clone());
}
for (i, wi) in words.iter().enumerate() {
for (j, wj) in words.iter().enumerate().skip(i + 1) {
if i / (SUBS * EACH) != j / (SUBS * EACH) {
assert!(
wi.iter().all(|w| !wj.contains(w)),
"descriptors {i} and {j} are from different groups and share a word"
);
}
}
}
let mut firsts: Vec<u32> = words.iter().map(|w| w[0]).collect();
firsts.sort_unstable();
firsts.dedup();
assert!(firsts.len() > GROUPS, "the second level separated nothing");
}
}
#[cfg(test)]
mod bench {
use super::*;
struct Lcg(u64);
impl Lcg {
fn byte(&mut self) -> u8 {
self.0 = self.0.wrapping_mul(6364136223846793005).wrapping_add(1442695040888963407);
(self.0 >> 33) as u8
}
}
#[test]
#[ignore]
fn quantise_timings() {
let mut rng = Lcg(0x1234_5678_9abc_def0);
let n = 60_000;
let desc: Vec<u8> = (0..n * DESC_LEN).map(|_| rng.byte()).collect();
let p = VocabParams { depth: 4, sample: 40_000, ..Default::default() };
let v = Vocabulary::build(&desc, &p);
let mut out = Vec::new();
let mut best = f64::MAX;
for _ in 0..7 {
let t = std::time::Instant::now();
for d in desc.chunks_exact(DESC_LEN) {
v.quantise(d, &mut out);
std::hint::black_box(&out);
}
best = best.min(t.elapsed().as_secs_f64() * 1e9 / n as f64);
}
println!("quantise: {best:8.1} ns/descriptor ({} words)", v.n_words());
}
#[test]
#[ignore]
fn shared_timings() {
let mut rng = Lcg(0x9e37_79b9_7f4a_7c15);
let mut w32 = || {
let a = rng.byte() as u32;
let b = rng.byte() as u32;
let c = rng.byte() as u32;
(a << 16) | (b << 8) | c
};
for (l, v) in [(1300usize, 1_771_561u32), (1070, 1_048_576), (1300, 262_144)] {
let lists: Vec<WordList> = (0..64)
.map(|_| {
let mut pairs: Vec<(u32, u32)> =
(0..l).map(|k| (w32() % v, (k / 3) as u32)).collect();
pairs.sort_unstable();
WordList {
word: pairs.iter().map(|p| p.0).collect(),
kp: pairs.iter().map(|p| p.1).collect(),
}
})
.collect();
let mut out = Vec::new();
let mut best = f64::MAX;
for _ in 0..7 {
let t = std::time::Instant::now();
let mut n = 0usize;
for a in lists.iter() {
for b in lists.iter() {
shared(a, b, &mut out, 60_000);
std::hint::black_box(&out);
n += 1;
}
}
best = best.min(t.elapsed().as_secs_f64() * 1e9 / n as f64);
}
println!("shared: {best:9.0} ns/call (len {l}, {v} words)");
}
}
#[test]
#[ignore]
fn query_timings() {
for (n_imgs, words) in [(9285usize, 170_000u32), (5637, 160_000)] {
let mut rng = Lcg(0x8e3f_1a9c_2b7d_4e51);
let lists: Vec<WordList> = (0..n_imgs)
.map(|_| {
let mut pairs: Vec<(u32, u32)> = (0..1300)
.map(|k| {
let a = (rng.byte() as u32) << 16 | (rng.byte() as u32) << 8 | rng.byte() as u32;
let b = (rng.byte() as u32) << 8 | rng.byte() as u32;
let w = if b % 4 == 0 { a % (words / 64) } else { a % words };
(w, (k / 3) as u32)
})
.collect();
pairs.sort_unstable();
WordList {
word: pairs.iter().map(|p| p.0).collect(),
kp: pairs.iter().map(|p| p.1).collect(),
}
})
.collect();
let inv = InvertedFile::build(&lists, words as usize, (n_imgs / 5).max(32));
let mut acc = vec![0f32; n_imgs];
let (mut touched, mut out) = (Vec::new(), Vec::new());
let mut best = f64::MAX;
let mut hits = 0usize;
for _ in 0..5 {
let t = std::time::Instant::now();
hits = 0;
for (i, wl) in lists.iter().enumerate().step_by(7) {
inv.query(wl, i as u32, &mut acc, &mut touched, &mut out);
hits += out.len();
std::hint::black_box(&out);
}
let n = lists.len().div_ceil(7);
best = best.min(t.elapsed().as_secs_f64() * 1e6 / n as f64);
}
println!(
"query: {best:8.1} us/query ({n_imgs} images, {words} words, {} candidates scored)",
hits / lists.len().div_ceil(7)
);
}
}
#[test]
#[ignore]
fn shared_overlap() {
let mut rng = Lcg(0x51ed_270b_6efc_2f4d);
let l = 1300usize;
let v = 1_771_561u32;
for shared_words in [0usize, 10, 40, 100, 300, 800] {
let lists: Vec<(WordList, WordList)> = (0..32)
.map(|_| {
let mut wa: Vec<u32> = Vec::new();
let mut wb: Vec<u32> = Vec::new();
let w32 = |rng: &mut Lcg| {
let (a, b, c) = (rng.byte() as u32, rng.byte() as u32, rng.byte() as u32);
((a << 16) | (b << 8) | c) % v
};
let mut common = Vec::new();
for _ in 0..shared_words {
common.push(w32(&mut rng));
}
for w in common.iter() {
for _ in 0..1 + (rng.byte() % 3) {
wa.push(*w);
}
for _ in 0..1 + (rng.byte() % 3) {
wb.push(*w);
}
}
while wa.len() < l {
wa.push(w32(&mut rng));
}
while wb.len() < l {
wb.push(w32(&mut rng));
}
let mk = |mut w: Vec<u32>| {
w.sort_unstable();
let kp: Vec<u32> = (0..w.len() as u32).collect();
WordList { word: w, kp }
};
(mk(wa), mk(wb))
})
.collect();
let mut out = Vec::new();
let mut best = f64::MAX;
let mut emitted = 0usize;
for _ in 0..7 {
let t = std::time::Instant::now();
let mut n = 0usize;
emitted = 0;
for (a, b) in lists.iter() {
shared(a, b, &mut out, 60_000);
emitted += out.len();
std::hint::black_box(&out);
n += 1;
}
best = best.min(t.elapsed().as_secs_f64() * 1e9 / n as f64);
}
println!(
"shared_overlap: {best:9.0} ns/call ({shared_words} shared words, {} pairs out)",
emitted / 32
);
}
}
#[test]
#[ignore]
fn shared_pool() {
let mut rng = Lcg(0x9e37_79b9_7f4a_7c15);
let mut w32 = || {
let a = rng.byte() as u32;
let b = rng.byte() as u32;
let c = rng.byte() as u32;
(a << 16) | (b << 8) | c
};
let l = 1300usize;
let v = 1_771_561u32;
for n_lists in [64usize, 1024, 8192] {
let lists: Vec<WordList> = (0..n_lists)
.map(|_| {
let mut pairs: Vec<(u32, u32)> =
(0..l).map(|k| (w32() % v, (k / 3) as u32)).collect();
pairs.sort_unstable();
WordList {
word: pairs.iter().map(|p| p.0).collect(),
kp: pairs.iter().map(|p| p.1).collect(),
}
})
.collect();
let mb = n_lists as f64 * l as f64 * 8.0 / 1e6;
let mut out = Vec::new();
let mut best = f64::MAX;
for _ in 0..5 {
let t = std::time::Instant::now();
let mut n = 0usize;
for q in 0..n_lists {
for step in 1..17 {
let o = (q * 7 + step * 613) % n_lists;
shared(&lists[q], &lists[o], &mut out, 60_000);
std::hint::black_box(&out);
n += 1;
}
}
best = best.min(t.elapsed().as_secs_f64() * 1e9 / n as f64);
}
println!("shared_pool: {best:9.0} ns/call ({n_lists} lists, {mb:.0} MB)");
}
}
#[test]
#[ignore]
fn quantise_depth() {
let mut rng = Lcg(0x1234_5678_9abc_def0);
let n = 60_000;
let desc: Vec<u8> = (0..n * DESC_LEN).map(|_| rng.byte()).collect();
for depth in 1..=5usize {
let p = VocabParams { depth, branching: 16, sample: 40_000, ..Default::default() };
let v = Vocabulary::build(&desc, &p);
let live: usize = v.levels.iter().map(|l| l.len() / DESC_LEN).sum();
let mut out = Vec::new();
let mut best = f64::MAX;
for _ in 0..7 {
let t = std::time::Instant::now();
for d in desc.chunks_exact(DESC_LEN) {
v.quantise(d, &mut out);
std::hint::black_box(&out);
}
best = best.min(t.elapsed().as_secs_f64() * 1e9 / n as f64);
}
let dists = 16 + (depth - 1) * 3 * 16;
println!(
"depth {depth}: {best:8.1} ns/descriptor {:5.2} ns/child-distance ({live} live centres, {:.1} MB)",
best / dists as f64,
live as f64 * DESC_LEN as f64 / 1e6
);
}
}
#[test]
#[ignore]
fn quantise_threads() {
let mut rng = Lcg(0x1234_5678_9abc_def0);
let n = 120_000;
let desc: Vec<u8> = (0..n * DESC_LEN).map(|_| rng.byte()).collect();
for (depth, sample) in [(3usize, 160_000usize), (4, 160_000), (5, 160_000)] {
let p = VocabParams { depth, branching: 16, sample, ..Default::default() };
let v = Vocabulary::build(&desc, &p);
let live: usize = v.levels.iter().map(|l| l.len() / DESC_LEN).sum();
let mut best1 = f64::MAX;
let mut best8 = f64::MAX;
for _ in 0..3 {
let mut out = Vec::new();
let t = std::time::Instant::now();
for d in desc.chunks_exact(DESC_LEN) {
v.quantise(d, &mut out);
std::hint::black_box(&out);
}
best1 = best1.min(t.elapsed().as_secs_f64() * 1e9 / n as f64);
let t = std::time::Instant::now();
desc.par_chunks(DESC_LEN * 64).for_each(|blk| {
let mut out = Vec::new();
for d in blk.chunks_exact(DESC_LEN) {
v.quantise(d, &mut out);
std::hint::black_box(&out);
}
});
best8 = best8.min(t.elapsed().as_secs_f64() * 1e9 / n as f64);
}
println!(
"depth {depth}: 1 thread {best1:8.1} ns/desc, all threads {best8:8.1} ns/desc (x{:.2} per thread on 8, {live} centres, {:.0} MB)",
best8 * 8.0 / best1,
live as f64 * DESC_LEN as f64 / 1e6
);
}
}
#[test]
#[ignore]
fn quantise_branching() {
let mut rng = Lcg(0x1234_5678_9abc_def0);
let n = 60_000;
let desc: Vec<u8> = (0..n * DESC_LEN).map(|_| rng.byte()).collect();
for branching in [8usize, 9, 10, 11, 12, 13, 14, 15, 16] {
let p = VocabParams { depth: 4, branching, sample: 40_000, ..Default::default() };
let v = Vocabulary::build(&desc, &p);
let mut out = Vec::new();
let mut best = f64::MAX;
for _ in 0..7 {
let t = std::time::Instant::now();
for d in desc.chunks_exact(DESC_LEN) {
v.quantise(d, &mut out);
std::hint::black_box(&out);
}
best = best.min(t.elapsed().as_secs_f64() * 1e9 / n as f64);
}
println!("branching {branching:2}: {best:8.1} ns/descriptor ({} words)", v.n_words());
}
}
}