pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
if a.len() != b.len() {
tracing::warn!(
"cosine_similarity length mismatch: {} vs {}, returning 0.0",
a.len(),
b.len()
);
return 0.0;
}
let mut dot = 0.0f32;
let mut norm_a = 0.0f32;
let mut norm_b = 0.0f32;
for (&x, &y) in a.iter().zip(b.iter()) {
dot += x * y;
norm_a += x * x;
norm_b += y * y;
}
if !norm_a.is_finite() || !norm_b.is_finite() || norm_a < 1e-8 || norm_b < 1e-8 {
return 0.0;
}
let sim = dot / (norm_a.sqrt() * norm_b.sqrt());
if sim.is_finite() { sim } else { 0.0 }
}
pub fn l2_normalize(vec: &mut [f32]) {
let norm: f32 = vec.iter().map(|x| x * x).sum::<f32>().sqrt();
if !norm.is_finite() {
vec.fill(0.0);
} else if norm > 1e-8 {
for v in vec.iter_mut() {
*v /= norm;
}
}
}
pub fn cosine_similarity_f32_f64(a: &[f32], b: &[f64]) -> f32 {
if a.len() != b.len() {
tracing::warn!(
"cosine_similarity_f32_f64 length mismatch: {} vs {}, returning 0.0",
a.len(),
b.len()
);
return 0.0;
}
let mut dot = 0.0f32;
let mut norm_a = 0.0f32;
let mut norm_b = 0.0f32;
for (&x, &y) in a.iter().zip(b.iter()) {
let y = y as f32;
dot += x * y;
norm_a += x * x;
norm_b += y * y;
}
if !norm_a.is_finite() || !norm_b.is_finite() || norm_a < 1e-8 || norm_b < 1e-8 {
return 0.0;
}
let sim = dot / (norm_a.sqrt() * norm_b.sqrt());
if sim.is_finite() { sim } else { 0.0 }
}
pub fn mean_vector(vectors: &[Vec<f32>]) -> Option<Vec<f32>> {
if vectors.is_empty() {
return None;
}
let dim = vectors[0].len();
let mut sum = vec![0.0f32; dim];
for v in vectors {
for (s, &x) in sum.iter_mut().zip(v.iter()) {
*s += x;
}
}
let n = vectors.len() as f32;
for s in &mut sum {
*s /= n;
}
Some(sum)
}
pub(crate) fn pairwise_cosine_similarity_matrix(embeddings: &[Vec<f32>]) -> Vec<f32> {
let n = embeddings.len();
let mut m = vec![0.0f32; n * n];
for i in 0..n {
m[i * n + i] = 1.0;
for j in (i + 1)..n {
let s = cosine_similarity(&embeddings[i], &embeddings[j]);
m[i * n + j] = s;
m[j * n + i] = s;
}
}
m
}
pub(crate) fn normalized_mean_centroids(
embeddings: &[Vec<f32>],
indices: &[usize],
labels: &[usize],
k: usize,
) -> Vec<Vec<f32>> {
let dim = embeddings.first().map(Vec::len).unwrap_or(0);
let mut sums = vec![vec![0.0f32; dim]; k];
let mut counts = vec![0usize; k];
for (&i, &l) in indices.iter().zip(labels.iter()) {
if l >= k {
continue;
}
if let Some(emb) = embeddings.get(i) {
for (a, &x) in sums[l].iter_mut().zip(emb.iter()) {
*a += x;
}
counts[l] += 1;
}
}
for (c, &n) in sums.iter_mut().zip(counts.iter()) {
if n > 0 {
for v in c.iter_mut() {
*v /= n as f32;
}
}
l2_normalize(c);
}
sums
}
use crate::types::Segment;
pub fn merge_segments(segments: Vec<Segment>, max_gap_secs: f64) -> Vec<Segment> {
if segments.is_empty() {
return segments;
}
let mut merged = Vec::new();
let mut current = segments[0].clone();
let (mut conf_sum, mut conf_count) = match current.confidence {
Some(c) => (c, 1u32),
None => (0.0, 0u32),
};
for next in segments.into_iter().skip(1) {
if current.speaker == next.speaker && next.time.start - current.time.end <= max_gap_secs {
current.time.end = next.time.end;
if let Some(c) = next.confidence {
conf_sum += c;
conf_count += 1;
}
} else {
current.confidence = mean_confidence(conf_sum, conf_count);
merged.push(current);
current = next;
(conf_sum, conf_count) = match current.confidence {
Some(c) => (c, 1u32),
None => (0.0, 0u32),
};
}
}
current.confidence = mean_confidence(conf_sum, conf_count);
merged.push(current);
merged
}
fn mean_confidence(sum: f32, count: u32) -> Option<f32> {
if count > 0 {
Some(sum / count as f32)
} else {
None
}
}
#[cfg(any(test, feature = "onnx"))]
pub(crate) struct ObjectPool<T> {
items: std::sync::Mutex<Vec<T>>,
}
#[cfg(any(test, feature = "onnx"))]
pub(crate) struct PooledGuard<'a, T> {
item: Option<T>,
pool: &'a ObjectPool<T>,
}
#[cfg(any(test, feature = "onnx"))]
impl<T> ObjectPool<T> {
pub(crate) fn new(items: Vec<T>) -> Self {
Self {
items: std::sync::Mutex::new(items),
}
}
pub(crate) fn checkout(&self) -> PooledGuard<'_, T> {
loop {
{
let mut guard = self.items.lock().unwrap_or_else(|e| e.into_inner());
if let Some(item) = guard.pop() {
return PooledGuard {
item: Some(item),
pool: self,
};
}
}
std::thread::yield_now();
}
}
}
#[cfg(any(test, feature = "onnx"))]
impl<T> std::ops::Deref for PooledGuard<'_, T> {
type Target = T;
fn deref(&self) -> &T {
match self.item.as_ref() {
Some(item) => item,
None => unreachable!("pooled item missing before Drop"),
}
}
}
#[cfg(any(test, feature = "onnx"))]
impl<T> std::ops::DerefMut for PooledGuard<'_, T> {
fn deref_mut(&mut self) -> &mut T {
match self.item.as_mut() {
Some(item) => item,
None => unreachable!("pooled item missing before Drop"),
}
}
}
#[cfg(any(test, feature = "onnx"))]
impl<T> Drop for PooledGuard<'_, T> {
fn drop(&mut self) {
if let Some(item) = self.item.take() {
self.pool
.items
.lock()
.unwrap_or_else(|e| e.into_inner())
.push(item);
}
}
}
#[derive(Clone, Debug)]
pub(crate) struct XorShift64Star {
state: u64,
}
impl XorShift64Star {
pub(crate) fn new(seed: u64) -> Self {
Self {
state: if seed == 0 {
0x9E37_79B9_7F4A_7C15
} else {
seed
},
}
}
#[inline]
pub(crate) fn next_u64(&mut self) -> u64 {
let mut x = self.state;
x ^= x >> 12;
x ^= x << 25;
x ^= x >> 27;
self.state = x;
x.wrapping_mul(0x2545_F491_4F6C_DD1D)
}
pub(crate) fn f64(&mut self) -> f64 {
(self.next_u64() >> 11) as f64 * (1.0 / ((1u64 << 53) as f64))
}
pub(crate) fn usize(&mut self, upper: usize) -> usize {
debug_assert!(upper > 0);
let upper = upper as u64;
((self.next_u64() as u128 * upper as u128) >> 64) as usize
}
}
#[allow(clippy::unwrap_used)]
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_cosine_similarity_identical() {
let a = vec![1.0, 2.0, 3.0];
let b = vec![1.0, 2.0, 3.0];
let sim = cosine_similarity(&a, &b);
assert!((sim - 1.0).abs() < 1e-5);
}
#[test]
fn test_cosine_similarity_orthogonal() {
let a = vec![1.0, 0.0, 0.0];
let b = vec![0.0, 1.0, 0.0];
let sim = cosine_similarity(&a, &b);
assert!(sim.abs() < 1e-5);
}
#[test]
fn test_cosine_similarity_opposite() {
let a = vec![1.0, 2.0, 3.0];
let b = vec![-1.0, -2.0, -3.0];
let sim = cosine_similarity(&a, &b);
assert!((sim + 1.0).abs() < 1e-5);
}
#[test]
fn test_l2_normalize() {
let mut v = vec![3.0, 4.0];
l2_normalize(&mut v);
assert!((v[0] - 0.6).abs() < 1e-5);
assert!((v[1] - 0.8).abs() < 1e-5);
}
#[test]
fn test_mean_vector() {
let vectors = vec![vec![1.0, 2.0], vec![3.0, 4.0]];
let mean = mean_vector(&vectors).unwrap();
assert!((mean[0] - 2.0).abs() < 1e-5);
assert!((mean[1] - 3.0).abs() < 1e-5);
}
#[test]
fn cosine_similarity_nan_input_returns_finite_zero() {
let a = vec![f32::NAN, 1.0, 0.0];
let b = vec![1.0, 0.0, 0.0];
let sim = cosine_similarity(&a, &b);
assert!(sim.is_finite());
assert_eq!(sim, 0.0);
}
#[test]
fn cosine_similarity_inf_input_returns_finite_zero() {
let a = vec![f32::INFINITY, 1.0, 0.0];
let b = vec![1.0, 2.0, 3.0];
let sim = cosine_similarity(&a, &b);
assert!(sim.is_finite());
assert_eq!(sim, 0.0);
}
#[test]
fn cosine_similarity_f32_f64_nan_input_returns_finite_zero() {
let a = vec![f32::NAN, 1.0, 0.0];
let b = vec![1.0_f64, 0.0, 0.0];
let sim = cosine_similarity_f32_f64(&a, &b);
assert!(sim.is_finite());
assert_eq!(sim, 0.0);
}
#[test]
fn l2_normalize_nan_input_yields_finite_vector() {
let mut v = vec![f32::NAN, 1.0, 2.0];
l2_normalize(&mut v);
assert!(v.iter().all(|x| x.is_finite()));
assert!(v.iter().all(|&x| x == 0.0));
}
#[test]
fn l2_normalize_inf_input_yields_finite_vector() {
let mut v = vec![f32::INFINITY, 1.0, 2.0];
l2_normalize(&mut v);
assert!(v.iter().all(|x| x.is_finite()));
assert!(v.iter().all(|&x| x == 0.0));
}
#[test]
fn pairwise_cosine_similarity_matrix_is_symmetric_with_unit_diagonal() {
let embeddings = vec![
vec![1.0, 0.0, 0.0],
vec![0.9, 0.1, 0.0],
vec![0.0, 1.0, 0.0],
];
let n = embeddings.len();
let m = pairwise_cosine_similarity_matrix(&embeddings);
assert_eq!(m.len(), n * n);
for i in 0..n {
assert_eq!(m[i * n + i], 1.0);
for j in 0..n {
assert_eq!(m[i * n + j], m[j * n + i]);
if i != j {
assert_eq!(
m[i * n + j],
cosine_similarity(&embeddings[i], &embeddings[j])
);
}
}
}
}
#[test]
fn pairwise_cosine_similarity_matrix_empty() {
assert!(pairwise_cosine_similarity_matrix(&[]).is_empty());
}
#[test]
fn normalized_mean_centroids_basic() {
let embeddings = vec![
vec![1.0, 0.0],
vec![1.0, 0.0],
vec![0.0, 2.0],
vec![0.0, 1.0],
];
let indices = [0, 1, 2, 3];
let labels = [0, 0, 1, 1];
let centroids = normalized_mean_centroids(&embeddings, &indices, &labels, 2);
assert_eq!(centroids.len(), 2);
for c in ¢roids {
let norm: f32 = c.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!((norm - 1.0).abs() < 1e-5, "centroid not unit norm");
}
assert!((centroids[0][0] - 1.0).abs() < 1e-5);
assert!((centroids[1][1] - 1.0).abs() < 1e-5);
}
#[test]
fn normalized_mean_centroids_empty_slot_stays_zero() {
let embeddings = vec![vec![1.0, 0.0]];
let centroids = normalized_mean_centroids(&embeddings, &[0], &[0], 3);
assert_eq!(centroids.len(), 3);
assert_eq!(centroids[1], vec![0.0, 0.0]);
assert_eq!(centroids[2], vec![0.0, 0.0]);
}
#[test]
fn normalized_mean_centroids_skips_out_of_range() {
let embeddings = vec![vec![1.0, 0.0]];
let centroids = normalized_mean_centroids(&embeddings, &[0, 5], &[7, 0], 2);
assert_eq!(centroids.len(), 2);
assert_eq!(centroids[0], vec![0.0, 0.0]);
assert_eq!(centroids[1], vec![0.0, 0.0]);
}
fn seg(start: f64, end: f64, spk: u32, conf: Option<f32>) -> Segment {
Segment {
time: crate::types::TimeRange { start, end },
speaker: Some(crate::types::SpeakerId(spk)),
confidence: conf,
}
}
#[test]
fn merge_confidence_is_order_independent_mean() {
let segs = vec![
seg(0.0, 1.0, 0, Some(0.6)),
seg(1.0, 2.0, 0, Some(0.9)),
seg(2.0, 3.0, 0, Some(0.9)),
];
let merged = merge_segments(segs, 0.5);
assert_eq!(merged.len(), 1);
let c = merged[0].confidence.expect("merged run has confidence");
assert!(
(c - 0.8).abs() < 1e-6,
"expected arithmetic mean 0.8, got {c}"
);
}
#[test]
fn merge_confidence_ignores_none_no_poisoning() {
let segs = vec![
seg(0.0, 1.0, 0, None),
seg(1.0, 2.0, 0, Some(0.8)),
seg(2.0, 3.0, 0, Some(0.6)),
];
let merged = merge_segments(segs, 0.5);
assert_eq!(merged.len(), 1);
let c = merged[0]
.confidence
.expect("present values must yield a mean");
assert!((c - 0.7).abs() < 1e-6, "expected 0.7, got {c}");
}
#[test]
fn merge_confidence_all_none_stays_none() {
let segs = vec![seg(0.0, 1.0, 0, None), seg(1.0, 2.0, 0, None)];
let merged = merge_segments(segs, 0.5);
assert_eq!(merged.len(), 1);
assert!(merged[0].confidence.is_none());
}
}