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:
@@ -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;
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user