lean_ctx/core/embeddings/
pooling.rs1pub 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
40pub 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
64pub 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
74pub 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 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 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}