bge_m3_embedding_server/embedder/
math.rs1use std::collections::HashMap;
18
19use ndarray::ArrayView1;
20
21pub(super) const SPECIAL_TOKENS: [u32; 4] = [0, 1, 2, 3];
23
24pub(super) fn normalize_l2(vec: &mut [f32]) {
26 let norm: f32 = vec.iter().map(|x| x * x).sum::<f32>().sqrt();
27 if norm > 0.0 {
28 for x in vec.iter_mut() {
29 *x /= norm;
30 }
31 }
32}
33
34pub(super) fn sparse_project(hidden: &[f32], weight: &ArrayView1<f32>, bias: f32) -> f32 {
38 let hidden_view = ArrayView1::from(hidden);
39 (hidden_view.dot(weight) + bias).max(0.0)
40}
41
42pub(super) fn sparse_maxpool(ids: &[u32], mask: &[u32], scores: &[f32]) -> (Vec<usize>, Vec<f32>) {
47 let mut token_weights: HashMap<usize, f32> = HashMap::new();
48
49 for (j, &token_id) in ids.iter().enumerate() {
50 if mask[j] == 0 {
51 continue;
52 }
53 if SPECIAL_TOKENS.contains(&token_id) {
54 continue;
55 }
56 let score = scores[j];
57 if score > 0.0 {
58 token_weights
59 .entry(token_id as usize)
60 .and_modify(|w| *w = w.max(score))
61 .or_insert(score);
62 }
63 }
64
65 let mut indices: Vec<usize> = token_weights.keys().copied().collect();
66 indices.sort_unstable();
67 let values: Vec<f32> = indices.iter().map(|k| token_weights[k]).collect();
68 (indices, values)
69}
70
71pub(super) fn median_usize(values: &mut [usize]) -> usize {
76 if values.is_empty() {
77 return 0;
78 }
79 values.sort_unstable();
80 values[values.len() / 2]
81}
82
83#[derive(Debug, Clone, Copy, Default)]
90pub(super) struct SeqLenDistribution {
91 pub min: usize,
93 pub max: usize,
95 pub mean: usize,
97 pub p95: usize,
103}
104
105pub(super) fn seq_len_distribution(lens: &[usize]) -> SeqLenDistribution {
110 if lens.is_empty() {
111 return SeqLenDistribution::default();
112 }
113 let min = *lens.iter().min().expect("non-empty");
114 let max = *lens.iter().max().expect("non-empty");
115 let mean = lens.iter().sum::<usize>() / lens.len();
116 let mut sorted = lens.to_vec();
117 sorted.sort_unstable();
118 let p95_idx = (sorted.len() * 95) / 100;
119 let p95 = sorted[p95_idx.min(sorted.len() - 1)];
120 SeqLenDistribution {
121 min,
122 max,
123 mean,
124 p95,
125 }
126}
127
128#[cfg(test)]
129mod tests {
130 use super::seq_len_distribution;
131
132 #[test]
133 fn single_element() {
134 let d = seq_len_distribution(&[42]);
135 assert_eq!(d.min, 42);
136 assert_eq!(d.max, 42);
137 assert_eq!(d.mean, 42);
138 assert_eq!(d.p95, 42);
139 }
140
141 #[test]
142 fn empty_returns_zeros() {
143 let d = seq_len_distribution(&[]);
144 assert_eq!(d.min, 0);
145 assert_eq!(d.max, 0);
146 assert_eq!(d.mean, 0);
147 assert_eq!(d.p95, 0);
148 }
149
150 #[test]
151 fn uniform_batch() {
152 let lens: Vec<usize> = vec![100; 64];
153 let d = seq_len_distribution(&lens);
154 assert_eq!(d.min, 100);
155 assert_eq!(d.max, 100);
156 assert_eq!(d.mean, 100);
157 assert_eq!(d.p95, 100);
158 }
159
160 #[test]
161 fn ascending_sequence_p95() {
162 let lens: Vec<usize> = (1..=100).collect();
165 let d = seq_len_distribution(&lens);
166 assert_eq!(d.min, 1);
167 assert_eq!(d.max, 100);
168 assert_eq!(d.mean, 50); assert_eq!(d.p95, 96);
170 }
171
172 #[test]
173 fn two_elements() {
174 let d = seq_len_distribution(&[10, 200]);
176 assert_eq!(d.min, 10);
177 assert_eq!(d.max, 200);
178 assert_eq!(d.mean, 105);
179 assert_eq!(d.p95, 200);
180 }
181
182 #[test]
183 fn p95_does_not_panic_on_small_batches() {
184 for n in 1usize..=20 {
186 let lens: Vec<usize> = (1..=n).collect();
187 let d = seq_len_distribution(&lens);
188 assert!(d.p95 >= d.min, "p95 < min for n={n}");
190 assert!(d.p95 <= d.max, "p95 > max for n={n}");
191 }
192 }
193
194 #[test]
195 fn mean_truncated() {
196 let d = seq_len_distribution(&[2, 3]);
198 assert_eq!(d.mean, 2);
199 }
200}