Skip to main content

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}