extern crate alloc;
use alloc::vec;
use alloc::vec::Vec;
use crate::matching::{MatchResult, Matcher, TimeOffset, clamp_score, frames_per_sec_compatible};
use crate::neural::{NeuralEmbedding, NeuralFingerprint};
#[derive(Copy, Clone, Debug, PartialEq, Eq)]
pub enum Aggregation {
Global,
SlidingMax,
Dtw,
}
#[derive(Clone, Debug)]
pub struct NeuralMatchConfig {
pub min_cosine: f32,
pub aggregation: Aggregation,
pub assume_normalized: bool,
}
impl Default for NeuralMatchConfig {
fn default() -> Self {
Self {
min_cosine: 0.80,
aggregation: Aggregation::SlidingMax,
assume_normalized: true,
}
}
}
pub struct NeuralMatcher {
cfg: NeuralMatchConfig,
}
impl Matcher for NeuralMatcher {
type Fingerprint = NeuralFingerprint;
type Config = NeuralMatchConfig;
fn new(cfg: Self::Config) -> Self {
Self { cfg }
}
fn config(&self) -> &Self::Config {
&self.cfg
}
fn match_one(&self, query: &Self::Fingerprint, reference: &Self::Fingerprint) -> MatchResult {
if !frames_per_sec_compatible(query.frames_per_sec, reference.frames_per_sec) {
return MatchResult::NONE;
}
let q = &query.embeddings;
let r = &reference.embeddings;
if q.is_empty() || r.is_empty() {
return MatchResult::NONE;
}
let dim = query.embedding_dim;
if dim == 0
|| reference.embedding_dim != dim
|| q.iter().any(|e| e.vector.len() != dim)
|| r.iter().any(|e| e.vector.len() != dim)
{
return MatchResult::NONE;
}
let assume_norm = self.cfg.assume_normalized;
let fps = reference.frames_per_sec;
match self.cfg.aggregation {
Aggregation::Global => self.match_global(q, r, dim, assume_norm),
Aggregation::SlidingMax => self.match_sliding(q, r, assume_norm, fps),
Aggregation::Dtw => self.match_dtw(q, r, assume_norm),
}
}
}
impl NeuralMatcher {
fn match_global(
&self,
q: &[NeuralEmbedding],
r: &[NeuralEmbedding],
dim: usize,
_assume_norm: bool,
) -> MatchResult {
let qc = normalized_centroid(q, dim);
let rc = normalized_centroid(r, dim);
let cos = dot(&qc, &rc);
self.build(
cos,
1.0,
TimeOffset::ZERO,
q.len().min(r.len()) as u32,
)
}
fn match_sliding(
&self,
q: &[NeuralEmbedding],
r: &[NeuralEmbedding],
assume_norm: bool,
fps: f32,
) -> MatchResult {
let query_is_short = q.len() <= r.len();
let (short, long) = if query_is_short { (q, r) } else { (r, q) };
let m = short.len();
let n = long.len();
let mut best_cos = f32::NEG_INFINITY;
let mut best_j = 0usize;
let mut sum_means = 0.0_f32;
let mut count = 0u32;
for j in 0..=(n - m) {
let mut acc = 0.0_f32;
for i in 0..m {
acc += cosine(&short[i].vector, &long[j + i].vector, assume_norm);
}
let mean = acc / m as f32;
sum_means += mean;
count += 1;
if mean > best_cos {
best_cos = mean;
best_j = j;
}
}
let delta = if query_is_short {
best_j as i64
} else {
-(best_j as i64)
};
let prominence = if count > 1 {
let mean_others = (sum_means - best_cos) / (count - 1) as f32;
(best_cos + 1.0) / (mean_others + 1.0)
} else {
1.0
};
self.build(
best_cos,
prominence,
TimeOffset::from_frames(delta, fps),
m as u32,
)
}
fn match_dtw(
&self,
q: &[NeuralEmbedding],
r: &[NeuralEmbedding],
assume_norm: bool,
) -> MatchResult {
let m = q.len();
let n = r.len();
let mut prev = vec![f32::INFINITY; n + 1];
let mut curr = vec![f32::INFINITY; n + 1];
prev[0] = 0.0;
for i in 1..=m {
curr[0] = f32::INFINITY;
for j in 1..=n {
let dist = 1.0 - cosine(&q[i - 1].vector, &r[j - 1].vector, assume_norm);
let best_prev = prev[j].min(curr[j - 1]).min(prev[j - 1]);
curr[j] = dist + best_prev;
}
core::mem::swap(&mut prev, &mut curr);
}
let total_cost = prev[n];
let path_len = m.max(n) as f32;
let mean_dist = if path_len > 0.0 {
total_cost / path_len
} else {
2.0
};
let equiv_cos = 1.0 - mean_dist;
let time_scale = if n > 0 { m as f32 / n as f32 } else { 1.0 };
let mut result = self.build(equiv_cos, 1.0, TimeOffset::ZERO, m.min(n) as u32);
result.time_scale = time_scale;
result
}
fn build(&self, cos: f32, prominence: f32, offset: TimeOffset, votes: u32) -> MatchResult {
let cos = if cos.is_finite() { cos } else { 0.0 };
let prominence = if prominence.is_finite() && prominence >= 0.0 {
prominence
} else {
0.0
};
MatchResult {
is_match: cos >= self.cfg.min_cosine,
score: clamp_score(cos),
votes,
prominence,
offset,
time_scale: 1.0,
}
}
}
#[inline]
fn dot(a: &[f32], b: &[f32]) -> f32 {
use wide::f32x8;
debug_assert_eq!(a.len(), b.len());
let n = a.len();
let chunks = n / 8;
let tail_start = chunks * 8;
let mut acc = f32x8::ZERO;
for i in 0..chunks {
let off = i * 8;
let va = f32x8::new(
a[off..off + 8]
.try_into()
.expect("a chunk is exactly 8 elements: loop iterates n/8 complete chunks"),
);
let vb = f32x8::new(
b[off..off + 8]
.try_into()
.expect("b chunk is exactly 8 elements: loop iterates n/8 complete chunks"),
);
acc = va.mul_add(vb, acc);
}
let mut sum = acc.reduce_add();
for i in tail_start..n {
sum += a[i] * b[i];
}
sum
}
#[inline]
fn cosine(a: &[f32], b: &[f32], assume_norm: bool) -> f32 {
let d = dot(a, b);
if assume_norm {
return d;
}
let na = crate::neural::embedder::sumsq_wide(a).sqrt();
let nb = crate::neural::embedder::sumsq_wide(b).sqrt();
if na < 1e-12 || nb < 1e-12 {
0.0
} else {
d / (na * nb)
}
}
fn normalized_centroid(embs: &[NeuralEmbedding], dim: usize) -> Vec<f32> {
let mut c = vec![0.0_f32; dim];
for e in embs {
for (acc, &v) in c.iter_mut().zip(e.vector.iter()) {
*acc += v;
}
}
let inv_n = 1.0 / embs.len() as f32;
for v in c.iter_mut() {
*v *= inv_n;
}
let norm = dot(&c, &c).sqrt();
if norm > 1e-12 {
let inv = 1.0 / norm;
for v in c.iter_mut() {
*v *= inv;
}
}
c
}
#[cfg(test)]
mod tests {
use super::*;
use crate::TimestampMs;
fn emb(v: &[f32], t_ms: u64) -> NeuralEmbedding {
let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt().max(1e-12);
NeuralEmbedding {
vector: v.iter().map(|x| x / norm).collect(),
t_start: TimestampMs(t_ms),
}
}
fn fp(embs: Vec<NeuralEmbedding>, dim: usize) -> NeuralFingerprint {
NeuralFingerprint {
embeddings: embs,
embedding_dim: dim,
frames_per_sec: 1.0,
}
}
#[test]
fn config_defaults() {
let c = NeuralMatchConfig::default();
assert!((c.min_cosine - 0.80).abs() < 1e-6);
assert_eq!(c.aggregation, Aggregation::SlidingMax);
assert!(c.assume_normalized);
}
#[test]
fn empty_returns_none() {
let m = NeuralMatcher::new(NeuralMatchConfig::default());
let empty = fp(vec![], 3);
let one = fp(vec![emb(&[1.0, 0.0, 0.0], 0)], 3);
assert_eq!(m.match_one(&empty, &one), MatchResult::NONE);
assert_eq!(m.match_one(&one, &empty), MatchResult::NONE);
}
#[test]
fn dim_mismatch_returns_none() {
let m = NeuralMatcher::new(NeuralMatchConfig::default());
let a = fp(vec![emb(&[1.0, 0.0, 0.0], 0)], 3);
let b = fp(vec![emb(&[1.0, 0.0], 0)], 2);
assert_eq!(m.match_one(&a, &b), MatchResult::NONE);
}
#[test]
fn self_match_cosine_one_global() {
let m = NeuralMatcher::new(NeuralMatchConfig {
aggregation: Aggregation::Global,
..Default::default()
});
let f = fp(
vec![
emb(&[1.0, 2.0, 3.0], 0),
emb(&[0.5, 0.1, 0.9], 1000),
emb(&[0.2, 0.8, 0.4], 2000),
],
3,
);
let res = m.match_one(&f, &f);
assert!(res.is_match, "self-match must be positive");
assert!(res.score > 0.999, "self-match score ~1: {}", res.score);
}
#[test]
fn self_match_sliding() {
let m = NeuralMatcher::new(NeuralMatchConfig::default());
let f = fp(
vec![
emb(&[1.0, 0.0, 0.0], 0),
emb(&[0.0, 1.0, 0.0], 1000),
emb(&[0.0, 0.0, 1.0], 2000),
],
3,
);
let res = m.match_one(&f, &f);
assert!(res.is_match);
assert!(res.score > 0.999, "score {}", res.score);
assert_eq!(res.offset.frames, 0);
}
#[test]
fn orthogonal_no_match() {
let m = NeuralMatcher::new(NeuralMatchConfig::default());
let a = fp(
vec![emb(&[1.0, 0.0, 0.0], 0), emb(&[1.0, 0.0, 0.0], 1000)],
3,
);
let b = fp(
vec![emb(&[0.0, 1.0, 0.0], 0), emb(&[0.0, 1.0, 0.0], 1000)],
3,
);
let res = m.match_one(&a, &b);
assert!(!res.is_match, "orthogonal embeddings must not match");
assert!(
res.score < 0.5,
"orthogonal score should be low: {}",
res.score
);
}
#[test]
fn sliding_offset_recovery_positive() {
let m = NeuralMatcher::new(NeuralMatchConfig::default());
let a = emb(&[1.0, 0.0, 0.0, 0.0], 0);
let b = emb(&[0.0, 1.0, 0.0, 0.0], 0);
let c = emb(&[0.0, 0.0, 1.0, 0.0], 0);
let d = emb(&[0.0, 0.0, 0.0, 1.0], 0);
let reference = fp(vec![a.clone(), b.clone(), c.clone(), d.clone()], 4);
let query = fp(vec![c.clone(), d.clone()], 4);
let res = m.match_one(&query, &reference);
assert!(res.is_match, "subsequence must match");
assert_eq!(res.offset.frames, 2, "query starts at reference window 2");
}
#[test]
fn sliding_offset_recovery_negative() {
let m = NeuralMatcher::new(NeuralMatchConfig::default());
let a = emb(&[1.0, 0.0, 0.0, 0.0], 0);
let b = emb(&[0.0, 1.0, 0.0, 0.0], 0);
let c = emb(&[0.0, 0.0, 1.0, 0.0], 0);
let d = emb(&[0.0, 0.0, 0.0, 1.0], 0);
let query = fp(vec![a.clone(), b.clone(), c.clone(), d.clone()], 4);
let reference = fp(vec![c.clone(), d.clone()], 4);
let res = m.match_one(&query, &reference);
assert!(res.is_match);
assert_eq!(res.offset.frames, -2, "reference aligns at query window 2");
}
#[test]
fn dtw_self_match() {
let m = NeuralMatcher::new(NeuralMatchConfig {
aggregation: Aggregation::Dtw,
..Default::default()
});
let f = fp(
vec![
emb(&[1.0, 0.0, 0.0], 0),
emb(&[0.0, 1.0, 0.0], 1000),
emb(&[0.0, 0.0, 1.0], 2000),
],
3,
);
let res = m.match_one(&f, &f);
assert!(
res.is_match,
"DTW self-match must be positive: {}",
res.score
);
assert!(res.score > 0.999, "DTW self-match score ~1: {}", res.score);
}
#[test]
fn dtw_tolerates_time_stretch() {
let m = NeuralMatcher::new(NeuralMatchConfig {
aggregation: Aggregation::Dtw,
min_cosine: 0.9,
..Default::default()
});
let a = emb(&[1.0, 0.0, 0.0], 0);
let b = emb(&[0.0, 1.0, 0.0], 0);
let reference = fp(vec![a.clone(), b.clone()], 3);
let query = fp(vec![a.clone(), a.clone(), b.clone()], 3);
let res = m.match_one(&query, &reference);
assert!(
res.is_match,
"DTW should tolerate the repeat: {}",
res.score
);
}
#[test]
fn not_normalized_config_still_works() {
let m = NeuralMatcher::new(NeuralMatchConfig {
assume_normalized: false,
aggregation: Aggregation::SlidingMax,
..Default::default()
});
let raw = |v: &[f32]| NeuralEmbedding {
vector: v.to_vec(),
t_start: TimestampMs(0),
};
let a = fp(vec![raw(&[2.0, 0.0, 0.0]), raw(&[0.0, 3.0, 0.0])], 3);
let b = fp(vec![raw(&[5.0, 0.0, 0.0]), raw(&[0.0, 7.0, 0.0])], 3);
let res = m.match_one(&a, &b);
assert!(
res.is_match,
"co-directional vectors must match: {}",
res.score
);
assert!(res.score > 0.999);
}
#[test]
fn determinism() {
let m = NeuralMatcher::new(NeuralMatchConfig::default());
let a = fp(
vec![emb(&[0.3, 0.7, 0.1], 0), emb(&[0.9, 0.2, 0.4], 1000)],
3,
);
let b = fp(
vec![emb(&[0.5, 0.5, 0.2], 0), emb(&[0.1, 0.8, 0.6], 1000)],
3,
);
assert_eq!(m.match_one(&a, &b), m.match_one(&a, &b));
}
}