Regroup the library without stopping the window
Pressing Regroup on a real library did not come back. Clustering 1,813 faces is the textbook agglomeration — compute every pairwise cosine, then repeatedly scan all live group pairs, score each with average link, and merge the best — and the scan is inside the loop. Each merge rescans every surviving pair, and each score is recomputed from scratch over every cross pair. Some 1.6 million pair scores per merge, some 700 merges to do. Three changes, none of which alter the answer. Only above-threshold pairs can ever matter. An average that reaches the threshold must have at least one term at or above it, so two groups with no qualifying pair between them can never merge — not now, and not after any sequence of merges, since merging only adds terms. The new `neighbours` module produces exactly that sparse list: 7,875 pairs rather than 1.6 million on the reference library. It also means the n^2 matrix is never materialised, so memory goes from O(n^2) to O(edges) — 2.5 GB to a few hundred KB at 25,000 faces. Merges cannot cross components, so the connected components of that graph are independent problems: four hundred small agglomerations instead of one large one. Average link is additive — sum(A u B, C) = sum(A, C) + sum(B, C) — so a merged group's scores follow by addition. Kept as running (sum, count) per adjacent pair, a score costs one division instead of a nested loop, and a heap with lazy invalidation replaces the rescan. Measured on the reference library: 0.28s, release, for all 1,813 faces. An exact ANN index was tried and removed, and neighbours.rs records why so it is not rediscovered as a good idea. IVF with a triangle-inequality bound is exact and prunes beautifully on synthetic clusters; on real embeddings it prunes *nothing* — 946 of 946 cell pairs survive. Median pair angle is 88.5 degrees and the merge threshold is 66.2, so the bound needs cells of radius under ~10 degrees, but two photographs of the same person sit 36-60 degrees apart. No ball-based partition of a 512-d near-orthogonal space can be tight enough. So the scan stayed exhaustive and got an unrolled dot product and its blocks spread across cores instead. Correctness is held by keeping the old implementation as an oracle: three tests run both engines over the same population — plain, under co-occurrence and anchor constraints, and with a size-weighted calibration — and assert the clusters are identical. Determinism is asserted at a size where the threaded path is in play. 62 tests pass. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
+628
-110
@@ -24,11 +24,50 @@
|
||||
//! Single-link chains: one bad edge welds two identities together, and it is
|
||||
//! the documented way face clustering fails on families. Average link asks
|
||||
//! whether the *groups* are similar, which one outlier cannot force.
|
||||
//!
|
||||
//! # How it runs, and why the obvious way does not
|
||||
//!
|
||||
//! The first implementation of this was the textbook one: compute every
|
||||
//! pairwise cosine, then repeatedly scan all live group pairs, score each with
|
||||
//! average link, and merge the best. It is correct, it is twenty lines, and on
|
||||
//! a real library it does not finish.
|
||||
//!
|
||||
//! The reason is that the scan is inside the loop. Each merge rescans every
|
||||
//! surviving pair — `O(g²)` of them — and each score is recomputed from
|
||||
//! scratch over every cross pair, `O(|A|·|B|)`. With 1,813 faces that is
|
||||
//! roughly 1.6 million pair scores per merge and some 700 merges to do; the
|
||||
//! window simply stops responding, which is what a user reports as "Regroup is
|
||||
//! broken". At 25,000 faces it is not slow, it is impossible.
|
||||
//!
|
||||
//! Three changes, none of which alter the answer:
|
||||
//!
|
||||
//! **Only above-threshold pairs can ever matter.** An average that reaches the
|
||||
//! threshold must have at least one term at or above it, so two groups with no
|
||||
//! qualifying pair between them can never merge — not now and not after any
|
||||
//! sequence of merges, since merging only adds terms. [`crate::neighbours`]
|
||||
//! produces exactly that sparse pair list, and everything below works on it.
|
||||
//! On the library above it is 7,875 pairs rather than 1.6 million.
|
||||
//!
|
||||
//! **Merges cannot cross components.** Groups only ever merge along those
|
||||
//! pairs, so the connected components of that graph are independent problems.
|
||||
//! A library of four hundred people becomes four hundred small agglomerations
|
||||
//! instead of one large one, and the quadratic term is paid per component.
|
||||
//!
|
||||
//! **Average link is additive.** `sum(A ∪ B, C) = sum(A, C) + sum(B, C)`, so a
|
||||
//! merged group's scores follow from the two it came from by addition — the
|
||||
//! Lance-Williams update. Kept as running `(sum, count)` per adjacent pair, a
|
||||
//! score costs one division instead of a nested loop, and a binary heap with
|
||||
//! lazy invalidation replaces the rescan.
|
||||
//!
|
||||
//! The output is unchanged, deliberately and testably so: `the_fast_engine_
|
||||
//! agrees_with_the_reference` runs both over the same population and asserts
|
||||
//! the clusters are identical.
|
||||
|
||||
use std::collections::HashSet;
|
||||
use std::cmp::Ordering;
|
||||
use std::collections::{BinaryHeap, HashMap, HashSet};
|
||||
|
||||
use crate::calibrate::Calibration;
|
||||
use crate::embedding::EMBEDDING_DIM;
|
||||
use crate::neighbours::{self, Faces};
|
||||
|
||||
/// Probability above which two groups are judged the same person.
|
||||
///
|
||||
@@ -80,79 +119,24 @@ pub fn cluster(faces: &[Candidate], cal: &Calibration, min_probability: f32) ->
|
||||
return Vec::new();
|
||||
}
|
||||
|
||||
let n = faces.len();
|
||||
let mut groups: Vec<Group> = faces
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(i, f)| Group {
|
||||
members: vec![i],
|
||||
images: HashSet::from([f.image]),
|
||||
person: f.confirmed_person,
|
||||
alive: true,
|
||||
})
|
||||
.collect();
|
||||
let embeddings: Vec<Vec<f32>> = faces.iter().map(|f| f.embedding.clone()).collect();
|
||||
let crop_px: Vec<f32> = faces.iter().map(|f| f.crop_px).collect();
|
||||
let images: Vec<u64> = faces.iter().map(|f| f.image).collect();
|
||||
let view = Faces {
|
||||
embeddings: &embeddings,
|
||||
crop_px: &crop_px,
|
||||
images: &images,
|
||||
};
|
||||
|
||||
// Pairwise cosine once. n² f32 is the honest cost at library scale — for
|
||||
// 25,000 faces that is the blocked GEMM docs/faces.md §9 describes, and the
|
||||
// caller is expected to shard rather than this function growing an index.
|
||||
let mut cos = vec![0.0_f32; n * n];
|
||||
for i in 0..n {
|
||||
for j in i + 1..n {
|
||||
let c = dot(&faces[i].embedding, &faces[j].embedding);
|
||||
cos[i * n + j] = c;
|
||||
cos[j * n + i] = c;
|
||||
}
|
||||
// Every pair that could ever contribute to a merge. See the module note on
|
||||
// why nothing outside this list can matter.
|
||||
let pairs = neighbours::above_threshold(&view, cal, min_probability);
|
||||
|
||||
let mut engine = Engine::new(faces, cal, min_probability);
|
||||
for component in components(faces.len(), &pairs) {
|
||||
engine.agglomerate(&component, &pairs);
|
||||
}
|
||||
|
||||
loop {
|
||||
let mut best: Option<(f32, usize, usize)> = None;
|
||||
|
||||
for a in 0..n {
|
||||
if !groups[a].alive {
|
||||
continue;
|
||||
}
|
||||
for b in a + 1..n {
|
||||
if !groups[b].alive || !can_link(&groups[a], &groups[b]) {
|
||||
continue;
|
||||
}
|
||||
let p = average_link(&groups[a], &groups[b], faces, &cos, n, cal);
|
||||
if p >= min_probability && best.is_none_or(|(bp, _, _)| p > bp) {
|
||||
best = Some((p, a, b));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let Some((_, a, b)) = best else { break };
|
||||
let taken = std::mem::take(&mut groups[b]);
|
||||
groups[b].alive = false;
|
||||
let ga = &mut groups[a];
|
||||
ga.members.extend(taken.members);
|
||||
ga.images.extend(taken.images);
|
||||
// At most one side carries a person: `can_link` refuses a merge of two
|
||||
// groups anchored to different people, so this cannot silently discard
|
||||
// one of them.
|
||||
ga.person = ga.person.or(taken.person);
|
||||
}
|
||||
|
||||
let mut out: Vec<Cluster> = groups
|
||||
.into_iter()
|
||||
.filter(|g| g.alive)
|
||||
.map(|mut g| {
|
||||
g.members.sort_unstable();
|
||||
Cluster {
|
||||
members: g.members,
|
||||
person: g.person,
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
// Largest first: the People view shows the best-evidenced groups at the top.
|
||||
out.sort_by(|x, y| {
|
||||
y.members
|
||||
.len()
|
||||
.cmp(&x.members.len())
|
||||
.then(x.members[0].cmp(&y.members[0]))
|
||||
});
|
||||
out
|
||||
engine.finish()
|
||||
}
|
||||
|
||||
/// Split one person's faces into the groups a raised threshold separates them
|
||||
@@ -178,62 +162,379 @@ pub fn split(faces: &[Candidate], cal: &Calibration, min_probability: f32) -> Ve
|
||||
cluster(&anchorless, cal, min_probability)
|
||||
}
|
||||
|
||||
#[derive(Debug, Default)]
|
||||
// ── the merge engine ──────────────────────────────────────────────────────
|
||||
|
||||
/// One live group, identified throughout by the index of its lowest member.
|
||||
#[derive(Debug)]
|
||||
struct Group {
|
||||
/// Ascending, always — [`Engine::cross`] sums in this order, and a stable
|
||||
/// order is what makes the floating-point total reproducible.
|
||||
members: Vec<usize>,
|
||||
images: HashSet<u64>,
|
||||
person: Option<u64>,
|
||||
alive: bool,
|
||||
/// Bumped on every merge, so heap entries naming an older state can be
|
||||
/// recognised and dropped instead of acted on.
|
||||
version: u64,
|
||||
}
|
||||
|
||||
/// Whether two groups are allowed to merge at all, before similarity is asked.
|
||||
fn can_link(a: &Group, b: &Group) -> bool {
|
||||
// Two confirmations of different people. The user has said these are not
|
||||
// the same person, and no similarity overrides that.
|
||||
if let (Some(pa), Some(pb)) = (a.person, b.person) {
|
||||
if pa != pb {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
// Co-occurrence: a photograph containing a face from each group means the
|
||||
// two faces are in the same frame, so they are not the same person.
|
||||
a.images.is_disjoint(&b.images)
|
||||
/// Running average-link state for one adjacent pair of groups.
|
||||
///
|
||||
/// `sum` is over **every** cross pair, not only the above-threshold ones —
|
||||
/// average link is an average over all of them, and counting only the
|
||||
/// qualifying pairs would report a similarity no group actually has.
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
struct Link {
|
||||
sum: f64,
|
||||
count: f64,
|
||||
}
|
||||
|
||||
/// Mean calibrated probability over every cross-group pair.
|
||||
fn average_link(
|
||||
a: &Group,
|
||||
b: &Group,
|
||||
faces: &[Candidate],
|
||||
cos: &[f32],
|
||||
n: usize,
|
||||
cal: &Calibration,
|
||||
) -> f32 {
|
||||
let mut sum = 0.0;
|
||||
let mut count = 0.0;
|
||||
for &i in &a.members {
|
||||
for &j in &b.members {
|
||||
let min_crop = faces[i].crop_px.min(faces[j].crop_px);
|
||||
sum += cal.probability(cos[i * n + j], min_crop, 0.0);
|
||||
count += 1.0;
|
||||
impl Link {
|
||||
fn probability(&self) -> f32 {
|
||||
if self.count == 0.0 {
|
||||
0.0
|
||||
} else {
|
||||
(self.sum / self.count) as f32
|
||||
}
|
||||
}
|
||||
if count == 0.0 {
|
||||
0.0
|
||||
}
|
||||
|
||||
/// A candidate merge, waiting in the heap.
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
struct Pending {
|
||||
probability: f32,
|
||||
a: usize,
|
||||
b: usize,
|
||||
/// Group versions when this was pushed. A mismatch on pop means a merge
|
||||
/// has happened since and a fresher entry for this pair is already queued.
|
||||
va: u64,
|
||||
vb: u64,
|
||||
}
|
||||
|
||||
impl PartialEq for Pending {
|
||||
fn eq(&self, other: &Self) -> bool {
|
||||
self.cmp(other) == Ordering::Equal
|
||||
}
|
||||
}
|
||||
impl Eq for Pending {}
|
||||
impl PartialOrd for Pending {
|
||||
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
|
||||
Some(self.cmp(other))
|
||||
}
|
||||
}
|
||||
impl Ord for Pending {
|
||||
/// Greatest pops first, so: highest probability, and on a tie the lowest
|
||||
/// index pair. That tiebreak is not cosmetic — it is what the old
|
||||
/// ascending scan did, and it is the whole of the determinism guarantee.
|
||||
fn cmp(&self, other: &Self) -> Ordering {
|
||||
self.probability
|
||||
.total_cmp(&other.probability)
|
||||
.then_with(|| other.a.cmp(&self.a))
|
||||
.then_with(|| other.b.cmp(&self.b))
|
||||
}
|
||||
}
|
||||
|
||||
struct Engine<'a> {
|
||||
faces: &'a [Candidate],
|
||||
cal: &'a Calibration,
|
||||
min_probability: f32,
|
||||
groups: Vec<Group>,
|
||||
links: HashMap<(usize, usize), Link>,
|
||||
/// Adjacency, as group ids. Kept alongside `links` so a merge can find
|
||||
/// everything it has to update without scanning the whole map.
|
||||
adjacent: Vec<HashSet<usize>>,
|
||||
}
|
||||
|
||||
impl<'a> Engine<'a> {
|
||||
fn new(faces: &'a [Candidate], cal: &'a Calibration, min_probability: f32) -> Self {
|
||||
let groups = faces
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(i, f)| Group {
|
||||
members: vec![i],
|
||||
images: HashSet::from([f.image]),
|
||||
person: f.confirmed_person,
|
||||
alive: true,
|
||||
version: 0,
|
||||
})
|
||||
.collect();
|
||||
Self {
|
||||
faces,
|
||||
cal,
|
||||
min_probability,
|
||||
groups,
|
||||
links: HashMap::new(),
|
||||
adjacent: vec![HashSet::new(); faces.len()],
|
||||
}
|
||||
}
|
||||
|
||||
/// Agglomerate one connected component to exhaustion.
|
||||
fn agglomerate(&mut self, component: &[usize], pairs: &[neighbours::Pair]) {
|
||||
if component.len() < 2 {
|
||||
return;
|
||||
}
|
||||
let members: HashSet<usize> = component.iter().copied().collect();
|
||||
|
||||
let mut heap = BinaryHeap::new();
|
||||
for p in pairs.iter().filter(|p| members.contains(&p.i)) {
|
||||
self.links.insert(
|
||||
key(p.i, p.j),
|
||||
Link {
|
||||
sum: p.probability as f64,
|
||||
count: 1.0,
|
||||
},
|
||||
);
|
||||
self.adjacent[p.i].insert(p.j);
|
||||
self.adjacent[p.j].insert(p.i);
|
||||
heap.push(Pending {
|
||||
probability: p.probability,
|
||||
a: p.i.min(p.j),
|
||||
b: p.i.max(p.j),
|
||||
va: 0,
|
||||
vb: 0,
|
||||
});
|
||||
}
|
||||
|
||||
while let Some(top) = heap.pop() {
|
||||
let Pending {
|
||||
probability,
|
||||
a,
|
||||
b,
|
||||
va,
|
||||
vb,
|
||||
} = top;
|
||||
|
||||
// Stale: one side has merged since this was queued, and the
|
||||
// replacement entry is already in the heap.
|
||||
if !self.groups[a].alive
|
||||
|| !self.groups[b].alive
|
||||
|| self.groups[a].version != va
|
||||
|| self.groups[b].version != vb
|
||||
{
|
||||
continue;
|
||||
}
|
||||
// The heap is ordered by probability, so the first entry below the
|
||||
// bar means nothing left in this component can reach it.
|
||||
if probability < self.min_probability {
|
||||
break;
|
||||
}
|
||||
if !self.can_link(a, b) {
|
||||
// Never becomes possible again: images only accumulate and an
|
||||
// anchor is never given up, so drop the pair for good.
|
||||
self.unlink(a, b);
|
||||
continue;
|
||||
}
|
||||
self.merge(a, b, &mut heap);
|
||||
}
|
||||
}
|
||||
|
||||
/// Whether two groups are allowed to merge at all, before similarity is
|
||||
/// asked.
|
||||
fn can_link(&self, a: usize, b: usize) -> bool {
|
||||
let (ga, gb) = (&self.groups[a], &self.groups[b]);
|
||||
// Two confirmations of different people. The user has said these are
|
||||
// not the same person, and no similarity overrides that.
|
||||
if let (Some(pa), Some(pb)) = (ga.person, gb.person) {
|
||||
if pa != pb {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
// Co-occurrence: a photograph containing a face from each group means
|
||||
// the two faces are in the same frame, so they are not the same person.
|
||||
ga.images.is_disjoint(&gb.images)
|
||||
}
|
||||
|
||||
/// Fold `b` into `a` and re-score everything that touched either.
|
||||
fn merge(&mut self, a: usize, b: usize, heap: &mut BinaryHeap<Pending>) {
|
||||
// Sorted, and deduplicated by the set: `a` and `b` may share
|
||||
// neighbours, and each must be visited once. Sorting is what keeps the
|
||||
// floating-point sums identical from run to run.
|
||||
let mut touched: Vec<usize> = self.adjacent[a]
|
||||
.union(&self.adjacent[b])
|
||||
.copied()
|
||||
.filter(|&c| c != a && c != b && self.groups[c].alive)
|
||||
.collect();
|
||||
touched.sort_unstable();
|
||||
|
||||
// Take the pair sums before the groups change underneath them.
|
||||
let carried: Vec<(usize, Option<Link>, Option<Link>)> = touched
|
||||
.iter()
|
||||
.map(|&c| {
|
||||
(
|
||||
c,
|
||||
self.links.get(&key(a, c)).copied(),
|
||||
self.links.get(&key(b, c)).copied(),
|
||||
)
|
||||
})
|
||||
.collect();
|
||||
|
||||
// a's own members, before b's are folded in. A missing a-side sum has
|
||||
// to be computed over these and not over the merged list, or b's
|
||||
// contribution would be counted twice.
|
||||
let a_members = self.groups[a].members.clone();
|
||||
|
||||
// Absorb b into a.
|
||||
let taken = std::mem::replace(
|
||||
&mut self.groups[b],
|
||||
Group {
|
||||
members: Vec::new(),
|
||||
images: HashSet::new(),
|
||||
person: None,
|
||||
alive: false,
|
||||
version: 0,
|
||||
},
|
||||
);
|
||||
{
|
||||
let ga = &mut self.groups[a];
|
||||
ga.members.extend(taken.members.iter().copied());
|
||||
ga.members.sort_unstable();
|
||||
ga.images.extend(taken.images.iter().copied());
|
||||
// At most one side carries a person: `can_link` refuses a merge of
|
||||
// two groups anchored to different people, so this cannot silently
|
||||
// discard one of them.
|
||||
ga.person = ga.person.or(taken.person);
|
||||
ga.version += 1;
|
||||
}
|
||||
|
||||
// b's own links are gone with it.
|
||||
for c in self.adjacent[b].clone() {
|
||||
self.links.remove(&key(b, c));
|
||||
self.adjacent[c].remove(&b);
|
||||
}
|
||||
self.adjacent[b].clear();
|
||||
self.links.remove(&key(a, b));
|
||||
self.adjacent[a].remove(&b);
|
||||
|
||||
for (c, from_a, from_b) in carried {
|
||||
// Dropping a pair the constraints now forbid saves computing a
|
||||
// score for a merge that can never happen — which for a newly
|
||||
// adjacent side is a real cost, not a bookkeeping one.
|
||||
if !self.can_link(a, c) {
|
||||
self.unlink(a, c);
|
||||
continue;
|
||||
}
|
||||
// A side with no stored link was not adjacent before, so its cross
|
||||
// pairs were all below threshold and were never summed. They still
|
||||
// belong in the average, so they are computed now — once, after
|
||||
// which the additive update carries them forward.
|
||||
let from_a = from_a.unwrap_or_else(|| self.cross(&a_members, c));
|
||||
let from_b = from_b.unwrap_or_else(|| self.cross(&taken.members, c));
|
||||
let merged = Link {
|
||||
sum: from_a.sum + from_b.sum,
|
||||
count: from_a.count + from_b.count,
|
||||
};
|
||||
self.links.insert(key(a, c), merged);
|
||||
self.adjacent[a].insert(c);
|
||||
self.adjacent[c].insert(a);
|
||||
|
||||
heap.push(Pending {
|
||||
probability: merged.probability(),
|
||||
a: a.min(c),
|
||||
b: a.max(c),
|
||||
va: self.groups[a.min(c)].version,
|
||||
vb: self.groups[a.max(c)].version,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
/// Exact `(sum, count)` over every cross pair between a member list and a
|
||||
/// group.
|
||||
///
|
||||
/// The one place a dot product is still computed during agglomeration, and
|
||||
/// it happens only when two groups become adjacent through a third — at
|
||||
/// which point their sub-threshold pairs, never summed because they were
|
||||
/// never interesting, have to be accounted for.
|
||||
fn cross(&self, members: &[usize], group: usize) -> Link {
|
||||
let mut sum = 0.0_f64;
|
||||
let mut count = 0.0_f64;
|
||||
for &i in members {
|
||||
for &j in &self.groups[group].members {
|
||||
let cos = neighbours::dot(&self.faces[i].embedding, &self.faces[j].embedding);
|
||||
let min_crop = self.faces[i].crop_px.min(self.faces[j].crop_px);
|
||||
sum += self.cal.probability(cos, min_crop, 0.0) as f64;
|
||||
count += 1.0;
|
||||
}
|
||||
}
|
||||
Link { sum, count }
|
||||
}
|
||||
|
||||
fn unlink(&mut self, a: usize, b: usize) {
|
||||
self.links.remove(&key(a, b));
|
||||
self.adjacent[a].remove(&b);
|
||||
self.adjacent[b].remove(&a);
|
||||
}
|
||||
|
||||
fn finish(self) -> Vec<Cluster> {
|
||||
let mut out: Vec<Cluster> = self
|
||||
.groups
|
||||
.into_iter()
|
||||
.filter(|g| g.alive)
|
||||
.map(|g| Cluster {
|
||||
members: g.members,
|
||||
person: g.person,
|
||||
})
|
||||
.collect();
|
||||
// Largest first: the People view shows the best-evidenced groups at the
|
||||
// top.
|
||||
out.sort_by(|x, y| {
|
||||
y.members
|
||||
.len()
|
||||
.cmp(&x.members.len())
|
||||
.then(x.members[0].cmp(&y.members[0]))
|
||||
});
|
||||
out
|
||||
}
|
||||
}
|
||||
|
||||
fn key(a: usize, b: usize) -> (usize, usize) {
|
||||
if a < b {
|
||||
(a, b)
|
||||
} else {
|
||||
sum / count
|
||||
(b, a)
|
||||
}
|
||||
}
|
||||
|
||||
fn dot(a: &[f32], b: &[f32]) -> f32 {
|
||||
debug_assert_eq!(a.len(), EMBEDDING_DIM);
|
||||
debug_assert_eq!(b.len(), EMBEDDING_DIM);
|
||||
a.iter().zip(b).map(|(x, y)| x * y).sum()
|
||||
}
|
||||
/// Connected components of the above-threshold graph.
|
||||
///
|
||||
/// Faces in different components can never end up in one group, so each is a
|
||||
/// separate and much smaller agglomeration. Returned with the members of each
|
||||
/// component ascending, and the components themselves in order of their lowest
|
||||
/// member — the determinism the merge order inherits.
|
||||
fn components(n: usize, pairs: &[neighbours::Pair]) -> Vec<Vec<usize>> {
|
||||
let mut parent: Vec<usize> = (0..n).collect();
|
||||
|
||||
fn find(parent: &mut [usize], mut x: usize) -> usize {
|
||||
while parent[x] != x {
|
||||
// Path halving: keeps the tree flat without a second pass.
|
||||
parent[x] = parent[parent[x]];
|
||||
x = parent[x];
|
||||
}
|
||||
x
|
||||
}
|
||||
|
||||
for p in pairs {
|
||||
let (ra, rb) = (find(&mut parent, p.i), find(&mut parent, p.j));
|
||||
if ra != rb {
|
||||
// Lowest root wins, so the representative of a component is
|
||||
// reproducible rather than an artefact of union order.
|
||||
let (lo, hi) = if ra < rb { (ra, rb) } else { (rb, ra) };
|
||||
parent[hi] = lo;
|
||||
}
|
||||
}
|
||||
|
||||
let mut by_root: HashMap<usize, Vec<usize>> = HashMap::new();
|
||||
for i in 0..n {
|
||||
let r = find(&mut parent, i);
|
||||
by_root.entry(r).or_default().push(i);
|
||||
}
|
||||
let mut out: Vec<Vec<usize>> = by_root.into_values().filter(|c| c.len() > 1).collect();
|
||||
out.sort_unstable_by_key(|c| c[0]);
|
||||
out
|
||||
}
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::embedding::EMBEDDING_DIM;
|
||||
|
||||
/// An embedding a known cosine away from a base direction, built by mixing
|
||||
/// two orthogonal unit vectors. Lets a test state "these two faces are 0.7
|
||||
@@ -441,4 +742,221 @@ mod tests {
|
||||
let small = sized.probability(0.5, 40.0, 0.0);
|
||||
assert!(big > small, "big {big} should beat small {small}");
|
||||
}
|
||||
|
||||
// ── the fast engine against the obvious one ───────────────────────────
|
||||
|
||||
/// The original implementation, kept as the oracle.
|
||||
///
|
||||
/// Deliberately the naive version this module replaced: rescan every live
|
||||
/// pair, score it from scratch over all cross pairs, merge the best,
|
||||
/// repeat. It is the definition of the answer, and the only thing the
|
||||
/// rewrite was allowed to change is how long it takes to get there.
|
||||
fn reference(faces: &[Candidate], cal: &Calibration, min_probability: f32) -> Vec<Cluster> {
|
||||
#[derive(Clone)]
|
||||
struct G {
|
||||
members: Vec<usize>,
|
||||
images: HashSet<u64>,
|
||||
person: Option<u64>,
|
||||
alive: bool,
|
||||
}
|
||||
if faces.is_empty() {
|
||||
return Vec::new();
|
||||
}
|
||||
let n = faces.len();
|
||||
let mut groups: Vec<G> = faces
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(i, f)| G {
|
||||
members: vec![i],
|
||||
images: HashSet::from([f.image]),
|
||||
person: f.confirmed_person,
|
||||
alive: true,
|
||||
})
|
||||
.collect();
|
||||
|
||||
let mut cos = vec![0.0_f32; n * n];
|
||||
for i in 0..n {
|
||||
for j in i + 1..n {
|
||||
let c = neighbours::dot(&faces[i].embedding, &faces[j].embedding);
|
||||
cos[i * n + j] = c;
|
||||
cos[j * n + i] = c;
|
||||
}
|
||||
}
|
||||
let linkable = |a: &G, b: &G| {
|
||||
if let (Some(pa), Some(pb)) = (a.person, b.person) {
|
||||
if pa != pb {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
a.images.is_disjoint(&b.images)
|
||||
};
|
||||
let average = |a: &G, b: &G| {
|
||||
let mut sum = 0.0_f32;
|
||||
let mut count = 0.0_f32;
|
||||
for &i in &a.members {
|
||||
for &j in &b.members {
|
||||
let min_crop = faces[i].crop_px.min(faces[j].crop_px);
|
||||
sum += cal.probability(cos[i * n + j], min_crop, 0.0);
|
||||
count += 1.0;
|
||||
}
|
||||
}
|
||||
if count == 0.0 {
|
||||
0.0
|
||||
} else {
|
||||
sum / count
|
||||
}
|
||||
};
|
||||
|
||||
loop {
|
||||
let mut best: Option<(f32, usize, usize)> = None;
|
||||
for a in 0..n {
|
||||
if !groups[a].alive {
|
||||
continue;
|
||||
}
|
||||
for b in a + 1..n {
|
||||
if !groups[b].alive || !linkable(&groups[a], &groups[b]) {
|
||||
continue;
|
||||
}
|
||||
let p = average(&groups[a], &groups[b]);
|
||||
if p >= min_probability && best.is_none_or(|(bp, _, _)| p > bp) {
|
||||
best = Some((p, a, b));
|
||||
}
|
||||
}
|
||||
}
|
||||
let Some((_, a, b)) = best else { break };
|
||||
let taken = groups[b].clone();
|
||||
groups[b].alive = false;
|
||||
groups[a].members.extend(taken.members);
|
||||
groups[a].images.extend(taken.images);
|
||||
groups[a].person = groups[a].person.or(taken.person);
|
||||
}
|
||||
|
||||
let mut out: Vec<Cluster> = groups
|
||||
.into_iter()
|
||||
.filter(|g| g.alive)
|
||||
.map(|mut g| {
|
||||
g.members.sort_unstable();
|
||||
Cluster {
|
||||
members: g.members,
|
||||
person: g.person,
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
out.sort_by(|x, y| {
|
||||
y.members
|
||||
.len()
|
||||
.cmp(&x.members.len())
|
||||
.then(x.members[0].cmp(&y.members[0]))
|
||||
});
|
||||
out
|
||||
}
|
||||
|
||||
/// `people` identities of `per` faces, each face in its own photograph,
|
||||
/// spread either side of the threshold so the population has genuine
|
||||
/// near-misses rather than obvious answers.
|
||||
fn population(people: usize, per: usize) -> Vec<Candidate> {
|
||||
let mut out = Vec::new();
|
||||
let mut image = 0u64;
|
||||
for p in 0..people {
|
||||
for m in 0..per {
|
||||
// Walks down through the merge boundary as m grows, so some
|
||||
// members join their group and some do not.
|
||||
let cosine = 1.0 - (m as f32) * 0.035;
|
||||
out.push(Candidate {
|
||||
face: out.len() as u64,
|
||||
image,
|
||||
embedding: at_cosine(p, cosine),
|
||||
crop_px: 60.0 + ((out.len() % 11) as f32) * 25.0,
|
||||
confirmed_person: None,
|
||||
});
|
||||
image += 1;
|
||||
}
|
||||
}
|
||||
out
|
||||
}
|
||||
|
||||
/// The point of the rewrite: same clusters, less work. A disagreement here
|
||||
/// is the rewrite being wrong, not the reference being slow.
|
||||
#[test]
|
||||
fn the_fast_engine_agrees_with_the_reference() {
|
||||
let faces = population(40, 6);
|
||||
assert_eq!(
|
||||
cluster(&faces, &cal(), DEFAULT_MERGE_PROBABILITY),
|
||||
reference(&faces, &cal(), DEFAULT_MERGE_PROBABILITY),
|
||||
);
|
||||
}
|
||||
|
||||
/// The two structural constraints are the ones a sparse graph could
|
||||
/// plausibly break, so they get their own comparison with anchors and
|
||||
/// co-occurrence in play.
|
||||
#[test]
|
||||
fn the_fast_engine_agrees_with_the_reference_under_constraints() {
|
||||
let mut faces = population(30, 6);
|
||||
// Some faces share a photograph, so cannot-link has to propagate
|
||||
// through groups that formed for other reasons.
|
||||
for i in (0..faces.len()).step_by(7) {
|
||||
faces[i].image = 900 + (i as u64 % 4);
|
||||
}
|
||||
// And some carry confirmations, including two of different people that
|
||||
// must never be brought together.
|
||||
for (n, i) in (0..faces.len()).step_by(11).enumerate() {
|
||||
faces[i].confirmed_person = Some(1 + (n as u64 % 3));
|
||||
}
|
||||
assert_eq!(
|
||||
cluster(&faces, &cal(), DEFAULT_MERGE_PROBABILITY),
|
||||
reference(&faces, &cal(), DEFAULT_MERGE_PROBABILITY),
|
||||
);
|
||||
}
|
||||
|
||||
/// The size term makes the merge boundary depend on the pair, which is the
|
||||
/// case the sparse pre-filter has to be built carefully to preserve.
|
||||
#[test]
|
||||
fn the_fast_engine_agrees_with_the_reference_with_a_size_term() {
|
||||
let faces = population(30, 6);
|
||||
let sized = Calibration {
|
||||
w_size: 0.4,
|
||||
b: -30.0 * 0.35 - 0.4 * 7.0,
|
||||
..cal()
|
||||
};
|
||||
assert_eq!(
|
||||
cluster(&faces, &sized, DEFAULT_MERGE_PROBABILITY),
|
||||
reference(&faces, &sized, DEFAULT_MERGE_PROBABILITY),
|
||||
);
|
||||
}
|
||||
|
||||
/// Determinism has to hold at a size where the indexed neighbour search is
|
||||
/// in play, not just on the handful of faces the small cases use.
|
||||
#[test]
|
||||
fn clustering_is_deterministic_at_scale() {
|
||||
let faces = population(200, 6);
|
||||
assert!(
|
||||
faces.len() > 1024,
|
||||
"population is below the indexing cutoff"
|
||||
);
|
||||
assert_eq!(
|
||||
cluster(&faces, &cal(), DEFAULT_MERGE_PROBABILITY),
|
||||
cluster(&faces, &cal(), DEFAULT_MERGE_PROBABILITY),
|
||||
);
|
||||
}
|
||||
|
||||
/// A face that matches nobody is left alone rather than being swept into
|
||||
/// the nearest group, and costs nothing to establish — it is in no
|
||||
/// component at all.
|
||||
#[test]
|
||||
fn a_face_matching_nothing_stays_on_its_own() {
|
||||
let mut faces = population(5, 4);
|
||||
faces.push(Candidate {
|
||||
face: 999,
|
||||
image: 5_000,
|
||||
embedding: at_cosine(200, 1.0),
|
||||
crop_px: 150.0,
|
||||
confirmed_person: None,
|
||||
});
|
||||
let out = cluster(&faces, &cal(), DEFAULT_MERGE_PROBABILITY);
|
||||
let last = faces.len() - 1;
|
||||
assert!(
|
||||
out.iter().any(|c| c.members == vec![last]),
|
||||
"the outlier was absorbed"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -39,6 +39,7 @@ pub mod detect;
|
||||
pub mod embed;
|
||||
pub mod embedding;
|
||||
pub mod naming;
|
||||
pub mod neighbours;
|
||||
|
||||
pub use align::{warp, Aligned112, Similarity, ALIGNED_EDGE, ARCFACE_TEMPLATE};
|
||||
pub use calibrate::{Calibration, Pairs, ReliabilityBand};
|
||||
|
||||
@@ -0,0 +1,566 @@
|
||||
//! TRACES: FR-CULL-9 | FR-CULL-10
|
||||
//! Finding the face pairs that could possibly be the same person.
|
||||
//!
|
||||
//! [`crate::cluster`] used to begin by computing every pairwise cosine and
|
||||
//! holding the lot as an `n²` matrix of `f32`. At 1,800 faces that is 13 MB,
|
||||
//! which is why it survived; at 25,000 it is 2.5 GB, which is why it could not
|
||||
//! keep surviving.
|
||||
//!
|
||||
//! Almost all of that matrix is thrown away unread. Clustering only ever asks
|
||||
//! whether a pair is *above* the merge threshold, and in a real library the
|
||||
//! answer is no for well over 99% of pairs — the reference library's 1,813
|
||||
//! faces produced 7,875 qualifying pairs out of 1.6 million. So this module
|
||||
//! answers the only question that is actually asked — **which pairs clear the
|
||||
//! bar** — and returns that sparse list. Memory goes from `O(n²)` to `O(edges)`
|
||||
//! and the caller never has to hold a matrix at all.
|
||||
//!
|
||||
//! # Exact, not approximate
|
||||
//!
|
||||
//! The usual way to make this fast is an approximate nearest-neighbour index,
|
||||
//! which trades recall for speed: it *misses* some true neighbours, and a
|
||||
//! missed neighbour here is a face that silently never joins its person.
|
||||
//! Nothing surfaces that — the screen just quietly shows one person as two —
|
||||
//! so it is a poor trade for a feature whose whole job is to be trusted. This
|
||||
//! module is exact, and the clusters it produces are identical to those from a
|
||||
//! full scan.
|
||||
//!
|
||||
//! # An exact index was tried, measured, and removed
|
||||
//!
|
||||
//! Worth recording so it is not rediscovered as a good idea. The obvious exact
|
||||
//! index is IVF with a triangle-inequality bound: group the embeddings into
|
||||
//! cells, and skip a whole cell **pair** when the geometry proves no member of
|
||||
//! one can reach any member of the other. On the sphere,
|
||||
//!
|
||||
//! ```text
|
||||
//! angle(x, y) >= angle(c_P, c_Q) - radius(P) - radius(Q)
|
||||
//! ```
|
||||
//!
|
||||
//! so a cell pair is impossible when `cos` of that lower bound falls below the
|
||||
//! threshold. Exact, no recall loss, and it prunes beautifully on synthetic
|
||||
//! clusters.
|
||||
//!
|
||||
//! It prunes **nothing at all** on real face embeddings. Measured over the
|
||||
//! 1,813-face reference library, at √n = 43 cells:
|
||||
//!
|
||||
//! | quantity | measured |
|
||||
//! |---|---|
|
||||
//! | median pair angle | 88.5° (cosine 0.026) |
|
||||
//! | merge threshold | 66.2° (cosine 0.403) |
|
||||
//! | median cell radius | 80.4° |
|
||||
//! | median centroid separation | 85.0° |
|
||||
//! | cell pairs surviving the bound | **946 of 946 — 100%** |
|
||||
//!
|
||||
//! The arithmetic is not close. For the bound to exclude a typical cell pair it
|
||||
//! needs `radius(P) + radius(Q) < 85° - 66° = 19°`, so cells of radius under
|
||||
//! ~10°. But two photographs of the *same person* sit 36–60° apart, so even a
|
||||
//! perfect single-identity cell has a radius three times too large. No
|
||||
//! ball-based partition of this space can have cells tight enough for the
|
||||
//! inequality to bite — 512-d embeddings are near-orthogonal, and that is the
|
||||
//! curse of dimensionality doing exactly what it says.
|
||||
//!
|
||||
//! So the scan stayed exhaustive, and the effort went where it actually pays:
|
||||
//! not materialising the matrix, an unrolled dot product, and spreading the
|
||||
//! blocks across cores. That is `O(n²)` time and `O(edges)` memory, which for
|
||||
//! this problem is the honest answer.
|
||||
|
||||
use crate::calibrate::Calibration;
|
||||
|
||||
/// Rows of the similarity triangle handed to one thread at a time.
|
||||
///
|
||||
/// Small enough that the tail of the triangle divides evenly across cores —
|
||||
/// row `i` does `n - i` comparisons, so equal *row counts* are very unequal
|
||||
/// work — and large enough that the per-block overhead disappears.
|
||||
const BLOCK: usize = 64;
|
||||
|
||||
/// Below this many faces, do the whole thing on the calling thread.
|
||||
///
|
||||
/// Spawning threads for a set this small costs more than the scan.
|
||||
const THREADS_ABOVE: usize = 2048;
|
||||
|
||||
/// One face pair that clears the merge threshold.
|
||||
///
|
||||
/// `i < j` always, and the probability is carried because the caller would
|
||||
/// otherwise recompute the sigmoid it took a dot product to reach.
|
||||
#[derive(Debug, Clone, Copy, PartialEq)]
|
||||
pub struct Pair {
|
||||
pub i: usize,
|
||||
pub j: usize,
|
||||
pub probability: f32,
|
||||
}
|
||||
|
||||
/// What the caller has to tell us about each face.
|
||||
///
|
||||
/// Deliberately not [`crate::cluster::Candidate`]: this module has no business
|
||||
/// knowing what a person or a photograph is, and taking the three arrays it
|
||||
/// actually reads keeps it testable on bare vectors.
|
||||
pub struct Faces<'a> {
|
||||
/// L2-normalised, `EMBEDDING_DIM` long, one per face.
|
||||
pub embeddings: &'a [Vec<f32>],
|
||||
/// Source pixels across the aligned crop, for the calibration's size term.
|
||||
pub crop_px: &'a [f32],
|
||||
/// Which photograph each face came from. Two faces in one frame are not
|
||||
/// the same person, so those pairs are never returned (docs/faces.md §9).
|
||||
pub images: &'a [u64],
|
||||
}
|
||||
|
||||
/// Every pair whose calibrated probability reaches `min_probability`.
|
||||
///
|
||||
/// Excludes pairs from the same photograph, which the clusterer would refuse
|
||||
/// anyway — dropping them here keeps them out of the graph the caller builds
|
||||
/// and out of the connected components it derives from it.
|
||||
///
|
||||
/// Ordered by `(i, j)`, which is what the caller's determinism rests on.
|
||||
pub fn above_threshold(faces: &Faces, cal: &Calibration, min_probability: f32) -> Vec<Pair> {
|
||||
let n = faces.embeddings.len();
|
||||
if n < 2 {
|
||||
return Vec::new();
|
||||
}
|
||||
|
||||
// The loosest cosine that could clear the bar for *any* pair in the set.
|
||||
// Cheaper than the sigmoid by far, and it rejects almost everything.
|
||||
let tau = loosest_cosine(faces.crop_px, cal, min_probability);
|
||||
|
||||
let blocks: Vec<(usize, usize)> = (0..n)
|
||||
.step_by(BLOCK)
|
||||
.map(|start| (start, (start + BLOCK).min(n)))
|
||||
.collect();
|
||||
|
||||
if n < THREADS_ABOVE {
|
||||
let mut out = Vec::new();
|
||||
for &(from, to) in &blocks {
|
||||
scan_block(faces, cal, min_probability, tau, from, to, &mut out);
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
// One worker per core bar one. This runs on a background thread behind a
|
||||
// button the user pressed, and NFR-ARCH-2 puts it behind the UI: taking
|
||||
// every core would stall the window it is reporting progress to.
|
||||
let workers = std::thread::available_parallelism()
|
||||
.map(|p| p.get().saturating_sub(1).max(1))
|
||||
.unwrap_or(1)
|
||||
.min(blocks.len());
|
||||
|
||||
let next = std::sync::atomic::AtomicUsize::new(0);
|
||||
let mut parts: Vec<Vec<Vec<Pair>>> = std::thread::scope(|scope| {
|
||||
let handles: Vec<_> = (0..workers)
|
||||
.map(|_| {
|
||||
let next = &next;
|
||||
let blocks = &blocks;
|
||||
scope.spawn(move || {
|
||||
// Results stay tagged with their block index, so the order
|
||||
// of the output does not depend on which thread got there
|
||||
// first. Determinism is a promise this module keeps.
|
||||
let mut mine: Vec<(usize, Vec<Pair>)> = Vec::new();
|
||||
loop {
|
||||
let b = next.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
|
||||
let Some(&(from, to)) = blocks.get(b) else {
|
||||
break;
|
||||
};
|
||||
let mut out = Vec::new();
|
||||
scan_block(faces, cal, min_probability, tau, from, to, &mut out);
|
||||
mine.push((b, out));
|
||||
}
|
||||
mine
|
||||
})
|
||||
})
|
||||
.collect();
|
||||
|
||||
let mut slots: Vec<Vec<Vec<Pair>>> = vec![Vec::new(); blocks.len()];
|
||||
for h in handles {
|
||||
for (b, pairs) in h.join().unwrap_or_default() {
|
||||
slots[b].push(pairs);
|
||||
}
|
||||
}
|
||||
slots
|
||||
});
|
||||
|
||||
let mut out = Vec::new();
|
||||
for slot in &mut parts {
|
||||
for pairs in slot.drain(..) {
|
||||
out.extend(pairs);
|
||||
}
|
||||
}
|
||||
out
|
||||
}
|
||||
|
||||
/// Compare rows `from..to` against everything after them.
|
||||
///
|
||||
/// The upper triangle, split by rows. Row `i` only looks at `j > i`, so every
|
||||
/// unordered pair is visited exactly once and the emitted order is `(i, j)`
|
||||
/// ascending within the block.
|
||||
fn scan_block(
|
||||
faces: &Faces,
|
||||
cal: &Calibration,
|
||||
min_probability: f32,
|
||||
tau: f32,
|
||||
from: usize,
|
||||
to: usize,
|
||||
out: &mut Vec<Pair>,
|
||||
) {
|
||||
let n = faces.embeddings.len();
|
||||
for i in from..to {
|
||||
let a = &faces.embeddings[i];
|
||||
let crop_a = faces.crop_px[i];
|
||||
let image_a = faces.images[i];
|
||||
for j in i + 1..n {
|
||||
if image_a == faces.images[j] {
|
||||
continue;
|
||||
}
|
||||
let cos = dot(a, &faces.embeddings[j]);
|
||||
// The cheap rejection, and it takes well over 99% of pairs.
|
||||
if cos < tau {
|
||||
continue;
|
||||
}
|
||||
let probability = cal.probability(cos, crop_a.min(faces.crop_px[j]), 0.0);
|
||||
if probability >= min_probability {
|
||||
out.push(Pair { i, j, probability });
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// The lowest cosine that could yield `min_probability` for any pair in the set.
|
||||
///
|
||||
/// The calibration is `sigmoid(a·cos + b + w_size·log2(crop))`, so for a fixed
|
||||
/// size term the cosine boundary is exact. The size term is *not* fixed — it
|
||||
/// varies per pair with the smaller of the two faces — so the safe bound uses
|
||||
/// whichever face size pushes the boundary lowest: the largest face when
|
||||
/// `w_size` is positive, the smallest when it is negative.
|
||||
///
|
||||
/// Returns [`f32::NEG_INFINITY`] — a filter that rejects nothing — where no
|
||||
/// boundary exists: a non-positive steepness, for which probability does not
|
||||
/// increase with cosine, or a threshold at the ends of the sigmoid. Those are
|
||||
/// degenerate calibrations rather than impossible ones, and the right response
|
||||
/// is to stop pruning, not to guess.
|
||||
fn loosest_cosine(crop_px: &[f32], cal: &Calibration, min_probability: f32) -> f32 {
|
||||
if cal.a <= 0.0 || !(min_probability > 0.0 && min_probability < 1.0) {
|
||||
return f32::NEG_INFINITY;
|
||||
}
|
||||
|
||||
let (mut lo, mut hi) = (f32::INFINITY, 0.0_f32);
|
||||
for &c in crop_px {
|
||||
let c = c.max(1.0);
|
||||
lo = lo.min(c);
|
||||
hi = hi.max(c);
|
||||
}
|
||||
if !lo.is_finite() {
|
||||
return f32::NEG_INFINITY;
|
||||
}
|
||||
|
||||
let extreme = if cal.w_size >= 0.0 { hi } else { lo };
|
||||
let tau = cal.boundary_at(min_probability, extreme, 0.0);
|
||||
if tau.is_nan() {
|
||||
return f32::NEG_INFINITY;
|
||||
}
|
||||
// Cosines never exceed 1, so a boundary above it legitimately rejects
|
||||
// everything. Clamped rather than left free so the comparison stays cheap.
|
||||
tau.min(1.0)
|
||||
}
|
||||
|
||||
/// Cosine of two L2-normalised embeddings.
|
||||
///
|
||||
/// Eight accumulators rather than one. Floating-point addition is not
|
||||
/// associative, so the compiler may not re-associate a single running total and
|
||||
/// the loop serialises on the adder's latency; eight independent chains give it
|
||||
/// something to pipeline and vectorise. The order is fixed and identical on
|
||||
/// every run, which is what the caller's determinism needs — it is a different
|
||||
/// order from the naive sum, not a variable one.
|
||||
pub(crate) fn dot(a: &[f32], b: &[f32]) -> f32 {
|
||||
const LANES: usize = 8;
|
||||
let mut acc = [0.0_f32; LANES];
|
||||
let chunks = a.len() / LANES;
|
||||
|
||||
for c in 0..chunks {
|
||||
let base = c * LANES;
|
||||
for (l, slot) in acc.iter_mut().enumerate() {
|
||||
*slot += a[base + l] * b[base + l];
|
||||
}
|
||||
}
|
||||
let mut total =
|
||||
((acc[0] + acc[1]) + (acc[2] + acc[3])) + ((acc[4] + acc[5]) + (acc[6] + acc[7]));
|
||||
for k in chunks * LANES..a.len() {
|
||||
total += a[k] * b[k];
|
||||
}
|
||||
total
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
const DIM: usize = 64;
|
||||
|
||||
fn cal() -> Calibration {
|
||||
Calibration {
|
||||
a: 30.0,
|
||||
b: -30.0 * 0.35,
|
||||
w_size: 0.0,
|
||||
valid: true,
|
||||
positive_pairs: 1000,
|
||||
negative_pairs: 10_000,
|
||||
}
|
||||
}
|
||||
|
||||
/// A unit vector in a reproducible pseudo-random direction.
|
||||
///
|
||||
/// Hashed from the seed rather than drawn from an RNG, so a failure is
|
||||
/// reproducible from the test alone.
|
||||
fn vector(seed: u64) -> Vec<f32> {
|
||||
let mut s = seed.wrapping_mul(0x9E37_79B9_7F4A_7C15) | 1;
|
||||
let mut v = Vec::with_capacity(DIM);
|
||||
for _ in 0..DIM {
|
||||
s ^= s << 13;
|
||||
s ^= s >> 7;
|
||||
s ^= s << 17;
|
||||
v.push(((s >> 11) as f64 / (1u64 << 53) as f64) as f32 - 0.5);
|
||||
}
|
||||
normalise(v)
|
||||
}
|
||||
|
||||
fn normalise(mut v: Vec<f32>) -> Vec<f32> {
|
||||
let n = v.iter().map(|x| x * x).sum::<f32>().sqrt();
|
||||
for x in &mut v {
|
||||
*x /= n;
|
||||
}
|
||||
v
|
||||
}
|
||||
|
||||
/// A vector a known cosine away from `base`.
|
||||
fn near(base: &[f32], other: &[f32], cosine: f32) -> Vec<f32> {
|
||||
let d = dot(base, other);
|
||||
let mut perp: Vec<f32> = other.iter().zip(base).map(|(o, b)| o - d * b).collect();
|
||||
let n = perp.iter().map(|x| x * x).sum::<f32>().sqrt();
|
||||
for x in &mut perp {
|
||||
*x /= n;
|
||||
}
|
||||
let s = (1.0 - cosine * cosine).max(0.0).sqrt();
|
||||
normalise(
|
||||
base.iter()
|
||||
.zip(&perp)
|
||||
.map(|(b, p)| cosine * b + s * p)
|
||||
.collect(),
|
||||
)
|
||||
}
|
||||
|
||||
struct Set {
|
||||
embeddings: Vec<Vec<f32>>,
|
||||
crop_px: Vec<f32>,
|
||||
images: Vec<u64>,
|
||||
}
|
||||
|
||||
impl Set {
|
||||
fn faces(&self) -> Faces<'_> {
|
||||
Faces {
|
||||
embeddings: &self.embeddings,
|
||||
crop_px: &self.crop_px,
|
||||
images: &self.images,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// `groups` identities, `per` faces each, every face in its own photograph.
|
||||
fn population(groups: usize, per: usize, tightness: f32) -> Set {
|
||||
let mut embeddings = Vec::new();
|
||||
let mut images = Vec::new();
|
||||
let mut image = 0u64;
|
||||
for g in 0..groups {
|
||||
let base = vector(g as u64 + 1);
|
||||
let off = vector(g as u64 + 9_999);
|
||||
for m in 0..per {
|
||||
embeddings.push(if m == 0 {
|
||||
base.clone()
|
||||
} else {
|
||||
near(&base, &off, tightness)
|
||||
});
|
||||
images.push(image);
|
||||
image += 1;
|
||||
}
|
||||
}
|
||||
let crop_px = vec![150.0; embeddings.len()];
|
||||
Set {
|
||||
embeddings,
|
||||
crop_px,
|
||||
images,
|
||||
}
|
||||
}
|
||||
|
||||
/// The unpruned, unthreaded, unblocked definition of the answer.
|
||||
fn reference(faces: &Faces, cal: &Calibration, min_probability: f32) -> Vec<Pair> {
|
||||
let n = faces.embeddings.len();
|
||||
let mut out = Vec::new();
|
||||
for i in 0..n {
|
||||
for j in i + 1..n {
|
||||
if faces.images[i] == faces.images[j] {
|
||||
continue;
|
||||
}
|
||||
let cos: f32 = faces.embeddings[i]
|
||||
.iter()
|
||||
.zip(&faces.embeddings[j])
|
||||
.map(|(x, y)| x * y)
|
||||
.sum();
|
||||
let p = cal.probability(cos, faces.crop_px[i].min(faces.crop_px[j]), 0.0);
|
||||
if p >= min_probability {
|
||||
out.push(Pair {
|
||||
i,
|
||||
j,
|
||||
probability: p,
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
out
|
||||
}
|
||||
|
||||
fn same_pairs(a: &[Pair], b: &[Pair]) -> bool {
|
||||
a.len() == b.len() && a.iter().zip(b).all(|(x, y)| x.i == y.i && x.j == y.j)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn nothing_to_pair_is_no_pairs() {
|
||||
let s = population(1, 1, 0.9);
|
||||
assert!(above_threshold(&s.faces(), &cal(), 0.9).is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn the_blocked_scan_finds_exactly_what_the_reference_does() {
|
||||
let s = population(60, 8, 0.97);
|
||||
let f = s.faces();
|
||||
let got = above_threshold(&f, &cal(), 0.9);
|
||||
let want = reference(&f, &cal(), 0.9);
|
||||
assert!(!want.is_empty(), "the reference found nothing to check");
|
||||
assert!(same_pairs(&got, &want), "{} vs {}", got.len(), want.len());
|
||||
}
|
||||
|
||||
/// Past `THREADS_ABOVE` the work is split across cores and stitched back
|
||||
/// together, and the stitching is where an order bug would live.
|
||||
#[test]
|
||||
fn the_threaded_scan_finds_exactly_what_the_reference_does() {
|
||||
let s = population(300, 8, 0.97);
|
||||
assert!(
|
||||
s.embeddings.len() > THREADS_ABOVE,
|
||||
"population is below the threading cutoff"
|
||||
);
|
||||
let f = s.faces();
|
||||
let got = above_threshold(&f, &cal(), 0.9);
|
||||
let want = reference(&f, &cal(), 0.9);
|
||||
assert!(!want.is_empty());
|
||||
assert!(same_pairs(&got, &want), "{} vs {}", got.len(), want.len());
|
||||
}
|
||||
|
||||
/// The size term moves the cosine boundary per pair, so the pre-filter has
|
||||
/// to be built from the most permissive size in the set or it will drop a
|
||||
/// pair that would have qualified.
|
||||
#[test]
|
||||
fn a_size_weighted_calibration_still_matches_the_reference() {
|
||||
let mut s = population(60, 8, 0.97);
|
||||
for (i, c) in s.crop_px.iter_mut().enumerate() {
|
||||
*c = 40.0 + (i % 17) as f32 * 30.0;
|
||||
}
|
||||
let sized = Calibration {
|
||||
w_size: 0.5,
|
||||
b: -30.0 * 0.35 - 0.5 * 7.0,
|
||||
..cal()
|
||||
};
|
||||
let f = s.faces();
|
||||
assert!(same_pairs(
|
||||
&above_threshold(&f, &sized, 0.9),
|
||||
&reference(&f, &sized, 0.9)
|
||||
));
|
||||
}
|
||||
|
||||
/// A negative size weight flips which extreme is permissive. Cheap to get
|
||||
/// wrong and silent when it is, so it gets its own case.
|
||||
#[test]
|
||||
fn a_negative_size_weight_prunes_from_the_other_end() {
|
||||
let mut s = population(60, 8, 0.97);
|
||||
for (i, c) in s.crop_px.iter_mut().enumerate() {
|
||||
*c = 40.0 + (i % 17) as f32 * 30.0;
|
||||
}
|
||||
let sized = Calibration {
|
||||
w_size: -0.5,
|
||||
b: -30.0 * 0.35 + 0.5 * 7.0,
|
||||
..cal()
|
||||
};
|
||||
let f = s.faces();
|
||||
assert!(same_pairs(
|
||||
&above_threshold(&f, &sized, 0.9),
|
||||
&reference(&f, &sized, 0.9)
|
||||
));
|
||||
}
|
||||
|
||||
/// A degenerate calibration has no cosine boundary to prune against, and
|
||||
/// must stop pruning rather than prune on a bound that does not hold.
|
||||
#[test]
|
||||
fn a_flat_calibration_prunes_nothing_and_still_agrees() {
|
||||
let flat = Calibration {
|
||||
a: 0.0,
|
||||
b: 4.0,
|
||||
..cal()
|
||||
};
|
||||
assert_eq!(loosest_cosine(&[150.0], &flat, 0.9), f32::NEG_INFINITY);
|
||||
let s = population(20, 6, 0.97);
|
||||
let f = s.faces();
|
||||
assert!(same_pairs(
|
||||
&above_threshold(&f, &flat, 0.9),
|
||||
&reference(&f, &flat, 0.9)
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn two_faces_in_one_photograph_are_never_paired() {
|
||||
let mut s = population(1, 2, 1.0);
|
||||
s.images = vec![7, 7];
|
||||
assert!(above_threshold(&s.faces(), &cal(), 0.9).is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn pairs_come_back_in_index_order() {
|
||||
let s = population(300, 8, 0.97);
|
||||
let pairs = above_threshold(&s.faces(), &cal(), 0.9);
|
||||
assert!(pairs
|
||||
.windows(2)
|
||||
.all(|w| (w[0].i, w[0].j) < (w[1].i, w[1].j)));
|
||||
assert!(pairs.iter().all(|p| p.i < p.j));
|
||||
}
|
||||
|
||||
/// Threads must not make the answer depend on which one finished first.
|
||||
#[test]
|
||||
fn the_same_input_yields_the_same_pairs() {
|
||||
let s = population(300, 8, 0.97);
|
||||
assert_eq!(
|
||||
above_threshold(&s.faces(), &cal(), 0.9),
|
||||
above_threshold(&s.faces(), &cal(), 0.9)
|
||||
);
|
||||
}
|
||||
|
||||
/// The unrolled dot has to agree with the obvious one, tail included — the
|
||||
/// lengths here are deliberately not multiples of the lane count.
|
||||
#[test]
|
||||
fn the_unrolled_dot_matches_the_naive_one() {
|
||||
for len in [1usize, 7, 8, 9, 63, 64, 65, 512] {
|
||||
let a: Vec<f32> = (0..len).map(|i| (i as f32 * 0.37).sin()).collect();
|
||||
let b: Vec<f32> = (0..len).map(|i| (i as f32 * 0.11).cos()).collect();
|
||||
let naive: f32 = a.iter().zip(&b).map(|(x, y)| x * y).sum();
|
||||
assert!(
|
||||
(dot(&a, &b) - naive).abs() < 1e-4,
|
||||
"len {len}: {} vs {naive}",
|
||||
dot(&a, &b)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
/// The pre-filter is the whole speed story, so it is worth asserting it
|
||||
/// actually rejects the bulk of the population rather than trusting it to.
|
||||
#[test]
|
||||
fn the_threshold_filter_rejects_almost_everything() {
|
||||
let s = population(60, 8, 0.97);
|
||||
let n = s.embeddings.len();
|
||||
let total = n * (n - 1) / 2;
|
||||
let kept = above_threshold(&s.faces(), &cal(), 0.9).len();
|
||||
assert!(
|
||||
kept * 20 < total,
|
||||
"kept {kept} of {total} pairs, which is not sparse"
|
||||
);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user