Let the merge engine's dot product use the machine's kernel too

Engine::cross is the one place a dot product is computed during
agglomeration — when two groups become adjacent through a third and
their sub-threshold pairs, never summed because they were never
interesting, have to be accounted for. It was calling the portable loop
while the scan beside it had AVX2 or NEON, which on the reference
library was 1,753,514 dot products taking 0.54s.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-08-29 11:58:07 +02:00
co-authored by Claude Opus 5
parent b4d39ba33a
commit a67402961d
2 changed files with 8 additions and 3 deletions
+6 -1
View File
@@ -370,6 +370,10 @@ impl Ord for Pending {
struct Engine<'a> {
faces: &'a [Candidate],
cal: &'a Calibration,
/// The same kernel [`neighbours`] scans with. [`Engine::cross`] is the one
/// place a dot product is computed during agglomeration, and it was using
/// the portable loop while the scan beside it had the machine's SIMD.
dot: neighbours::DotFn,
min_probability: f32,
groups: Vec<Group>,
links: HashMap<(usize, usize), Link>,
@@ -394,6 +398,7 @@ impl<'a> Engine<'a> {
Self {
faces,
cal,
dot: neighbours::fastest_dot(),
min_probability,
groups,
links: HashMap::new(),
@@ -581,7 +586,7 @@ impl<'a> Engine<'a> {
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 cos = (self.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;
+2 -2
View File
@@ -340,7 +340,7 @@ fn loosest_cosine(crop_px: &[f32], cal: &Calibration, min_probability: f32) -> f
/// 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;
pub(crate) type DotFn = fn(&[f32], &[f32]) -> f32;
/// The widest dot product this machine can actually run.
///
@@ -369,7 +369,7 @@ type DotFn = fn(&[f32], &[f32]) -> f32;
/// 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 {
pub(crate) fn fastest_dot() -> DotFn {
#[cfg(target_arch = "x86_64")]
{
if std::arch::is_x86_feature_detected!("avx2") && std::arch::is_x86_feature_detected!("fma")