Skip to main content

lean_ctx/core/embeddings/
pooling.rs

1//! Pooling strategies for transformer hidden states.
2//!
3//! Converts per-token hidden states `[seq_len × dim]` into a single
4//! fixed-size embedding vector `[dim]`.
5
6/// Mean pooling over token positions, weighted by attention mask.
7///
8/// Takes the raw hidden state output `[1 × seq_len × dim]` flattened to a Vec,
9/// and produces a single embedding by averaging across attended positions.
10pub fn mean_pool(
11    hidden_states: &[f32],
12    attention_mask: &[i32],
13    seq_len: usize,
14    dim: usize,
15) -> Vec<f32> {
16    let mut sum = vec![0.0f32; dim];
17    let mut count = 0.0f32;
18
19    for pos in 0..seq_len {
20        if attention_mask.get(pos).copied().unwrap_or(0) > 0 {
21            let offset = pos * dim;
22            for (d, sum_val) in sum.iter_mut().enumerate().take(dim) {
23                if let Some(&val) = hidden_states.get(offset + d) {
24                    *sum_val += val;
25                }
26            }
27            count += 1.0;
28        }
29    }
30
31    if count > 0.0 {
32        for val in &mut sum {
33            *val /= count;
34        }
35    }
36
37    sum
38}
39
40/// Batched mean pooling over multiple sequences in a single output tensor.
41///
42/// The model output is `[batch, max_seq_len, dim]` flattened row-major. Each
43/// sequence is mean-pooled using its own attention mask to exclude padding.
44pub fn mean_pool_batch(
45    hidden_states: &[f32],
46    masks: &[&[i32]],
47    max_seq_len: usize,
48    dim: usize,
49) -> Vec<Vec<f32>> {
50    let batch = masks.len();
51    let expected_len = batch * max_seq_len * dim;
52    if hidden_states.len() < expected_len {
53        return vec![vec![0.0; dim]; batch];
54    }
55    let mut results = Vec::with_capacity(batch);
56    for (b, m) in masks.iter().enumerate().take(batch) {
57        let offset = b * max_seq_len * dim;
58        let h = &hidden_states[offset..][..max_seq_len * dim];
59        results.push(mean_pool(h, m, max_seq_len, dim));
60    }
61    results
62}
63
64/// L2-normalize a vector in-place.
65pub fn normalize_l2(vec: &mut [f32]) {
66    let norm: f32 = vec.iter().map(|x| x * x).sum::<f32>().sqrt();
67    if norm > f32::EPSILON {
68        for x in vec.iter_mut() {
69            *x /= norm;
70        }
71    }
72}
73
74/// Compute the L2 norm of a vector.
75pub fn l2_norm(vec: &[f32]) -> f32 {
76    vec.iter().map(|x| x * x).sum::<f32>().sqrt()
77}
78
79#[cfg(test)]
80mod tests {
81    use super::*;
82
83    #[test]
84    fn mean_pool_basic() {
85        // 2 tokens, 3 dimensions, all attended
86        let hidden = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
87        let mask = vec![1, 1];
88        let result = mean_pool(&hidden, &mask, 2, 3);
89        assert_eq!(result.len(), 3);
90        assert!((result[0] - 2.5).abs() < 1e-6);
91        assert!((result[1] - 3.5).abs() < 1e-6);
92        assert!((result[2] - 4.5).abs() < 1e-6);
93    }
94
95    #[test]
96    fn mean_pool_with_padding() {
97        // 3 tokens, 2 dimensions, last token is padding
98        let hidden = vec![1.0, 2.0, 3.0, 4.0, 0.0, 0.0];
99        let mask = vec![1, 1, 0];
100        let result = mean_pool(&hidden, &mask, 3, 2);
101        assert!((result[0] - 2.0).abs() < 1e-6);
102        assert!((result[1] - 3.0).abs() < 1e-6);
103    }
104
105    #[test]
106    fn mean_pool_single_token() {
107        let hidden = vec![5.0, 10.0];
108        let mask = vec![1];
109        let result = mean_pool(&hidden, &mask, 1, 2);
110        assert!((result[0] - 5.0).abs() < 1e-6);
111        assert!((result[1] - 10.0).abs() < 1e-6);
112    }
113
114    #[test]
115    fn mean_pool_all_masked() {
116        let hidden = vec![1.0, 2.0, 3.0, 4.0];
117        let mask = vec![0, 0];
118        let result = mean_pool(&hidden, &mask, 2, 2);
119        assert!(result.iter().all(|&v| v == 0.0));
120    }
121
122    #[test]
123    fn normalize_l2_basic() {
124        let mut vec = vec![3.0, 4.0];
125        normalize_l2(&mut vec);
126        assert!((vec[0] - 0.6).abs() < 1e-6);
127        assert!((vec[1] - 0.8).abs() < 1e-6);
128
129        let norm: f32 = vec.iter().map(|x| x * x).sum::<f32>().sqrt();
130        assert!((norm - 1.0).abs() < 1e-5);
131    }
132
133    #[test]
134    fn normalize_l2_already_normalized() {
135        let mut vec = vec![1.0, 0.0, 0.0];
136        normalize_l2(&mut vec);
137        assert!((vec[0] - 1.0).abs() < 1e-6);
138    }
139
140    #[test]
141    fn normalize_l2_zero_vector() {
142        let mut vec = vec![0.0, 0.0, 0.0];
143        normalize_l2(&mut vec);
144        assert!(vec.iter().all(|&v| v == 0.0));
145    }
146
147    #[test]
148    fn l2_norm_basic() {
149        assert!((l2_norm(&[3.0, 4.0]) - 5.0).abs() < 1e-6);
150    }
151
152    #[test]
153    fn l2_norm_unit() {
154        assert!((l2_norm(&[1.0, 0.0, 0.0]) - 1.0).abs() < 1e-6);
155    }
156}