const MIN_NORM: f32 = 1e-12;
pub(crate) fn mean_pool_normalise<'a>(
acc: &mut [f32],
rows: impl Iterator<Item = &'a [f32]>,
) -> usize {
debug_assert!(acc.iter().all(|x| *x == 0.0));
let mut count = 0usize;
for row in rows {
for (a, w) in acc.iter_mut().zip(row) {
*a += *w;
}
count += 1;
}
if count == 0 {
return 0;
}
let inv = 1.0 / count as f32;
for a in acc.iter_mut() {
*a *= inv;
}
if !normalise(acc) {
acc.fill(0.0);
return 0;
}
count
}
pub(crate) fn normalise(v: &mut [f32]) -> bool {
let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if !(norm.is_finite() && norm > MIN_NORM) {
return false;
}
let inv = 1.0 / norm;
for x in v.iter_mut() {
*x *= inv;
}
true
}
pub fn is_zero(v: &[f32]) -> bool {
v.iter().all(|x| *x == 0.0)
}