Give the similarity scan the machine's SIMD, and its cache

The scan is O(n²) dot products and nothing else, so its speed is the
face subsystem's speed — and it was running at 0.7 flops per cycle.

Two separate faults, both measured over the reference 18,143-face
library on twenty cores. It walked the whole embedding array once per
row, ~336 GB of traffic, where a column tile that fits in L2 is read
once per tile of rows: 4.64s → 2.81s. And the workspace builds for
baseline x86-64 — SSE2, no FMA — into which the portable loop was not
being vectorised at all: 2.81s → 0.86s, 195 GFLOP/s.

So the dot product is now chosen per machine. AVX2 + FMA where
is_x86_feature_detected! finds it; NEON unconditionally on aarch64,
since Advanced SIMD is in that baseline and every Android device the app
builds for has it — with the explicit vfmaq, because LLVM will not fuse
a multiply and an add without being told to. The portable loop stays as
the definition the others are tested against, and
the_fastest_kernel_agrees_with_the_portable_one is the only check the
NEON path gets on a machine that is not aarch64.

Faces::embeddings is one flat buffer rather than a Vec per face: the
pointer chase defeated both the prefetcher and the tiling, and it is
also the layout a GPU pass would want.

Behaviour is unchanged and that is checked rather than asserted — the
same 1,531,969 pairs from all three kernels, and on the real library the
same 2,518 groups holding the same 16,246 faces with the same confidence
distribution. A full regroup there goes from 10.0s to 5.9s; the rest is
the agglomeration, which is a sequential heap walk and is where the next
look should go.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-08-29 11:33:15 +02:00
co-authored by Claude Opus 5
parent ebb7d3cf5c
commit e596eb0657
4 changed files with 321 additions and 43 deletions
+276 -40
View File
@@ -72,6 +72,13 @@ use crate::calibrate::Calibration;
/// work — and large enough that the per-block overhead disappears.
const BLOCK: usize = 64;
/// Columns compared against one row block before moving on.
///
/// The other half of the tiling: 64 embeddings of 512 floats is 128 KB, which
/// sits in L2 beside the row block instead of being re-read from memory for
/// every row. See [`scan_rows`] for what it was worth.
const COLUMN_TILE: 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.
@@ -94,8 +101,17 @@ pub struct Pair {
/// 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>],
/// Every embedding end to end, L2-normalised, [`Faces::dim`] floats each.
///
/// **Flat, not a slice of vectors**, and the difference is measurable: a
/// `&[Vec<f32>]` is one heap allocation per face and the scan chases a
/// pointer per row, which defeats both the prefetcher and the tiling this
/// module does to stay in cache. One buffer is also the layout a GPU would
/// want, which is where this is eventually going.
pub embeddings: &'a [f32],
/// Floats per embedding — [`crate::EMBEDDING_DIM`] in practice, a parameter
/// so the tests can work in 64 dimensions.
pub dim: usize,
/// 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
@@ -103,6 +119,26 @@ pub struct Faces<'a> {
pub images: &'a [u64],
}
impl Faces<'_> {
/// How many faces there are.
pub fn len(&self) -> usize {
if self.dim == 0 {
0
} else {
self.embeddings.len() / self.dim
}
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
#[inline(always)]
fn row(&self, i: usize) -> &[f32] {
&self.embeddings[i * self.dim..(i + 1) * self.dim]
}
}
/// Every pair whose calibrated probability reaches `min_probability`.
///
/// Excludes pairs from the same photograph, which the clusterer would refuse
@@ -111,14 +147,20 @@ pub struct Faces<'a> {
///
/// 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();
let n = faces.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 scan = Scan {
faces,
cal,
min_probability,
tau: loosest_cosine(faces.crop_px, cal, min_probability),
dot: fastest_dot(),
};
let blocks: Vec<(usize, usize)> = (0..n)
.step_by(BLOCK)
@@ -128,7 +170,7 @@ pub fn above_threshold(faces: &Faces, cal: &Calibration, min_probability: f32) -
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);
scan_rows(&scan, from, to, &mut out);
}
return out;
}
@@ -158,7 +200,7 @@ pub fn above_threshold(faces: &Faces, cal: &Calibration, min_probability: f32) -
break;
};
let mut out = Vec::new();
scan_block(faces, cal, min_probability, tau, from, to, &mut out);
scan_rows(&scan, from, to, &mut out);
mine.push((b, out));
}
mine
@@ -184,40 +226,69 @@ pub fn above_threshold(faces: &Faces, cal: &Calibration, min_probability: f32) -
out
}
/// Compare rows `from..to` against everything after them.
/// Everything [`scan_rows`] needs that does not change between blocks.
///
/// 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,
/// A struct rather than eight arguments, and the grouping is real: these five
/// are fixed for a whole scan and only the row range moves.
#[derive(Clone, Copy)]
struct Scan<'a> {
faces: &'a Faces<'a>,
cal: &'a 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] {
dot: DotFn,
}
/// Compare rows `from..to` against everything after them.
///
/// The upper triangle, split by rows, and **tiled on the column side too**.
/// Walking `j` from `i + 1` to `n` for one row at a time streams the whole
/// embedding array past the core once per row — 336 GB of traffic for an
/// 18,000-face library — where a column tile small enough to sit in L2 is read
/// once per *tile* of rows. Measured on that library, the tiling alone took the
/// scan from 4.64 s to 2.81 s before any change to the kernel.
///
/// Pairs come out ordered by `(i, j)`: the tiles are walked in ascending order
/// but the row loop is inside them, so the block's own output is sorted before
/// it is returned. Blocks are concatenated in row order, so the whole list is
/// ordered — which is the promise the caller's determinism rests on.
fn scan_rows(scan: &Scan, from: usize, to: usize, out: &mut Vec<Pair>) {
let Scan {
faces,
cal,
min_probability,
tau,
dot,
} = *scan;
let n = faces.len();
for tile in (from..n).step_by(COLUMN_TILE) {
let tile_end = (tile + COLUMN_TILE).min(n);
for i in from..to {
// The diagonal: nothing at or before `i` is this row's business.
let start = tile.max(i + 1);
if start >= tile_end {
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 });
let a = faces.row(i);
let crop_a = faces.crop_px[i];
let image_a = faces.images[i];
for j in start..tile_end {
if image_a == faces.images[j] {
continue;
}
let cos = dot(a, faces.row(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 });
}
}
}
}
out.sort_unstable_by_key(|p| (p.i, p.j));
}
/// The lowest cosine that could yield `min_probability` for any pair in the set.
@@ -266,10 +337,144 @@ fn loosest_cosine(crop_px: &[f32], cal: &Calibration, min_probability: f32) -> f
/// 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.
/// A dot product over two equal-length, L2-normalised rows.
///
/// Chosen once per scan rather than per pair — see [`fastest_dot`].
type DotFn = fn(&[f32], &[f32]) -> f32;
/// The widest dot product this machine can actually run.
///
/// # Why this is worth unsafe code
///
/// The scan is the arithmetic floor of the whole subsystem and it was running
/// at **0.7 flops per cycle**. The workspace builds for baseline `x86-64`,
/// which is SSE2 and no FMA, and the portable loop below was not being
/// vectorised into even that. Measured over a real 18,143-face library, on
/// twenty cores:
///
/// | kernel | scan | GFLOP/s |
/// |---|---|---|
/// | portable, untiled (what this replaced) | 4.64 s | 36 |
/// | portable, tiled | 2.81 s | 60 |
/// | **AVX2 + FMA, tiled** | **0.86 s** | **195 |
///
/// Identical pair lists — 1,531,969 — all three ways.
///
/// **Runtime detection on x86-64, unconditional on aarch64.** Advanced SIMD is
/// in the aarch64 baseline, so every Android device that runs this has NEON and
/// there is nothing to detect; on x86-64 AVX2 is not baseline and a binary that
/// assumed it would not start on older hardware.
///
/// The three kernels sum in different orders, so a cosine may differ in its
/// last bit between them. That only matters for a pair sitting exactly on the
/// threshold, and the portable kernel already sums in eight accumulators rather
/// than one, so the module was never bit-comparable with a naive sum.
fn fastest_dot() -> DotFn {
#[cfg(target_arch = "x86_64")]
{
if std::arch::is_x86_feature_detected!("avx2") && std::arch::is_x86_feature_detected!("fma")
{
return dot_avx2;
}
}
#[cfg(target_arch = "aarch64")]
{
return dot_neon;
}
#[allow(unreachable_code)]
dot
}
/// Eight-wide fused multiply-add, on the half of desktops that have it.
#[cfg(target_arch = "x86_64")]
fn dot_avx2(a: &[f32], b: &[f32]) -> f32 {
// SAFETY: `fastest_dot` is the only thing that hands this out, and only
// after `is_x86_feature_detected!` has said both features are present.
unsafe { dot_avx2_inner(a, b) }
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2", enable = "fma")]
unsafe fn dot_avx2_inner(a: &[f32], b: &[f32]) -> f32 {
use std::arch::x86_64::*;
let n = a.len().min(b.len());
let (mut acc0, mut acc1) = (_mm256_setzero_ps(), _mm256_setzero_ps());
let mut k = 0;
// Two accumulators, because one FMA cannot start until the previous one
// retires and the unit is pipelined several deep.
while k + 16 <= n {
// SAFETY: `k + 16 <= n`, and `n` is within both slices.
acc0 = _mm256_fmadd_ps(
_mm256_loadu_ps(a.as_ptr().add(k)),
_mm256_loadu_ps(b.as_ptr().add(k)),
acc0,
);
acc1 = _mm256_fmadd_ps(
_mm256_loadu_ps(a.as_ptr().add(k + 8)),
_mm256_loadu_ps(b.as_ptr().add(k + 8)),
acc1,
);
k += 16;
}
let mut lanes = [0.0_f32; 8];
_mm256_storeu_ps(lanes.as_mut_ptr(), _mm256_add_ps(acc0, acc1));
let mut total = ((lanes[0] + lanes[1]) + (lanes[2] + lanes[3]))
+ ((lanes[4] + lanes[5]) + (lanes[6] + lanes[7]));
for k in k..n {
total += a[k] * b[k];
}
total
}
/// The same, four-wide, for the phone and the tablet.
///
/// No feature detection and no `target_feature`: Advanced SIMD is mandatory in
/// the aarch64 baseline, so this compiles for every Android target the app
/// builds for. The explicit `vfmaq` matters — LLVM will not fuse a multiply and
/// an add on its own without fast-math, which is most of the win.
#[cfg(target_arch = "aarch64")]
fn dot_neon(a: &[f32], b: &[f32]) -> f32 {
use std::arch::aarch64::*;
let n = a.len().min(b.len());
// SAFETY: every load below is bounded by `k + 16 <= n`, and `n` is within
// both slices. NEON needs no feature detection on aarch64.
unsafe {
let mut acc = [vdupq_n_f32(0.0); 4];
let mut k = 0;
while k + 16 <= n {
for (l, slot) in acc.iter_mut().enumerate() {
*slot = vfmaq_f32(
*slot,
vld1q_f32(a.as_ptr().add(k + l * 4)),
vld1q_f32(b.as_ptr().add(k + l * 4)),
);
}
k += 16;
}
let mut total = vaddvq_f32(vaddq_f32(
vaddq_f32(acc[0], acc[1]),
vaddq_f32(acc[2], acc[3]),
));
for k in k..n {
total += a[k] * b[k];
}
total
}
}
/// The portable fallback, and the definition the others have to agree with.
///
/// Eight accumulators so the adds are independent; that is as much as can be
/// asked of a loop that has to compile for anything.
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;
let n = a.len().min(b.len());
let chunks = n / LANES;
for c in 0..chunks {
let base = c * LANES;
@@ -279,7 +484,7 @@ pub(crate) fn dot(a: &[f32], b: &[f32]) -> f32 {
}
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() {
for k in chunks * LANES..n {
total += a[k] * b[k];
}
total
@@ -344,7 +549,7 @@ mod tests {
}
struct Set {
embeddings: Vec<Vec<f32>>,
embeddings: Vec<f32>,
crop_px: Vec<f32>,
images: Vec<u64>,
}
@@ -353,10 +558,15 @@ mod tests {
fn faces(&self) -> Faces<'_> {
Faces {
embeddings: &self.embeddings,
dim: DIM,
crop_px: &self.crop_px,
images: &self.images,
}
}
fn len(&self) -> usize {
self.embeddings.len() / DIM
}
}
/// `groups` identities, `per` faces each, every face in its own photograph.
@@ -379,7 +589,7 @@ mod tests {
}
let crop_px = vec![150.0; embeddings.len()];
Set {
embeddings,
embeddings: embeddings.concat(),
crop_px,
images,
}
@@ -387,16 +597,17 @@ mod tests {
/// 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 n = faces.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]
let cos: f32 = faces
.row(i)
.iter()
.zip(&faces.embeddings[j])
.zip(faces.row(j))
.map(|(x, y)| x * y)
.sum();
let p = cal.probability(cos, faces.crop_px[i].min(faces.crop_px[j]), 0.0);
@@ -416,6 +627,31 @@ mod tests {
a.len() == b.len() && a.iter().zip(b).all(|(x, y)| x.i == y.i && x.j == y.j)
}
/// The guard on the SIMD kernels, and the only check the aarch64 one gets
/// on a machine that is not aarch64: whatever [`fastest_dot`] picked has to
/// agree with the portable definition. A wrong lane index or a mishandled
/// tail would show up here as a wildly different number, not a rounding
/// difference.
#[test]
fn the_fastest_kernel_agrees_with_the_portable_one() {
let fast = fastest_dot();
for seed in 0..64u64 {
let a = vector(seed);
let b = vector(seed + 1_000);
let (want, got) = (dot(&a, &b), fast(&a, &b));
assert!(
(want - got).abs() < 1e-5,
"kernel disagreed on seed {seed}: {want} vs {got}"
);
}
// A length that is not a multiple of the widest step, so the tail is
// exercised rather than assumed away.
let a: Vec<f32> = (0..37).map(|k| k as f32 * 0.01).collect();
let b: Vec<f32> = (0..37).map(|k| 1.0 - k as f32 * 0.02).collect();
assert!((dot(&a, &b) - fast(&a, &b)).abs() < 1e-5, "tail mishandled");
}
#[test]
fn nothing_to_pair_is_no_pairs() {
let s = population(1, 1, 0.9);
@@ -438,7 +674,7 @@ mod tests {
fn the_threaded_scan_finds_exactly_what_the_reference_does() {
let s = population(300, 8, 0.97);
assert!(
s.embeddings.len() > THREADS_ABOVE,
s.len() > THREADS_ABOVE,
"population is below the threading cutoff"
);
let f = s.faces();