cyberbrain_embed/pool.rs
1//! Mean pooling and L2 normalisation. Pure functions so they can be tested without a model.
2
3/// Below this norm a pooled vector is treated as empty rather than scaled up into noise.
4const MIN_NORM: f32 = 1e-12;
5
6/// Sums the rows named by `ids` (already filtered to known ids) into `acc`, divides by the
7/// count and L2-normalises in place. Returns the number of rows pooled.
8///
9/// If no row was pooled, or the mean is (numerically) the zero vector, `acc` is left all
10/// zero and `0` is returned. The zero vector is the documented output for "nothing to
11/// embed"; it is never NaN.
12pub(crate) fn mean_pool_normalise<'a>(
13 acc: &mut [f32],
14 rows: impl Iterator<Item = &'a [f32]>,
15) -> usize {
16 debug_assert!(acc.iter().all(|x| *x == 0.0));
17 let mut count = 0usize;
18 for row in rows {
19 for (a, w) in acc.iter_mut().zip(row) {
20 *a += *w;
21 }
22 count += 1;
23 }
24 if count == 0 {
25 return 0;
26 }
27 let inv = 1.0 / count as f32;
28 for a in acc.iter_mut() {
29 *a *= inv;
30 }
31 if !normalise(acc) {
32 acc.fill(0.0);
33 return 0;
34 }
35 count
36}
37
38/// Scales `v` to unit L2 norm in place. Returns `false`, leaving `v` untouched, when the
39/// norm is too small or not finite to normalise meaningfully.
40pub(crate) fn normalise(v: &mut [f32]) -> bool {
41 let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
42 if !(norm.is_finite() && norm > MIN_NORM) {
43 return false;
44 }
45 let inv = 1.0 / norm;
46 for x in v.iter_mut() {
47 *x *= inv;
48 }
49 true
50}
51
52/// True when every component is exactly zero: the vector this crate returns for an empty
53/// or all-unknown input. Callers should skip semantic search for such a query instead of
54/// ranking on cosines that are all 0.
55pub fn is_zero(v: &[f32]) -> bool {
56 v.iter().all(|x| *x == 0.0)
57}