1use thiserror::Error;
34
35#[derive(Debug, Clone, Copy, PartialEq, Eq)]
37pub enum PoolingType {
38 None,
40 Mean,
42 Cls,
46 Last,
48 Rank,
50}
51
52#[derive(Debug, Error)]
53pub enum PoolingError {
54 #[error(
55 "{key} = {value} is not one of llama.cpp's llama_pooling_type values \
56 (-1 unspecified, 0 NONE, 1 MEAN, 2 CLS, 3 LAST, 4 RANK)"
57 )]
58 UnknownWireValue { key: String, value: i64 },
59 #[error("{key} is present but is not an integer: {value}")]
60 NotAnInteger { key: String, value: String },
61 #[error(
62 "pooling type RANK is a reranker classification head (cls / cls_out / \
63 classifier.output_labels), not a pooling rule: pooling sees hidden states and a \
64 width, and cannot reach those matrices. POST /v1/rerank with a query and \
65 documents, which runs the head this checkpoint carries. Refusing rather than \
66 returning a CLS row that is not what RANK means"
67 )]
68 Unimplemented,
69 #[error("cannot pool an empty sequence")]
70 EmptySequence,
71 #[error("hidden states are {len} floats, which is not a whole number of {n_embd}-wide rows")]
72 RaggedHiddenStates { len: usize, n_embd: usize },
73}
74
75impl PoolingType {
76 pub fn name(self) -> &'static str {
79 match self {
80 PoolingType::None => "NONE",
81 PoolingType::Mean => "MEAN",
82 PoolingType::Cls => "CLS",
83 PoolingType::Last => "LAST",
84 PoolingType::Rank => "RANK",
85 }
86 }
87
88 fn from_wire(key: &str, value: i64) -> Result<Option<Self>, PoolingError> {
91 Ok(match value {
92 -1 => None,
93 0 => Some(PoolingType::None),
94 1 => Some(PoolingType::Mean),
95 2 => Some(PoolingType::Cls),
96 3 => Some(PoolingType::Last),
97 4 => Some(PoolingType::Rank),
98 other => {
99 return Err(PoolingError::UnknownWireValue {
100 key: key.to_string(),
101 value: other,
102 })
103 }
104 })
105 }
106
107 pub fn from_gguf(
111 file: &impl ferrox_gguf::TensorSource,
112 arch: &str,
113 ) -> Result<Option<Self>, PoolingError> {
114 let key = format!("{arch}.pooling_type");
115 let Some(value) = file.metadata(&key) else {
116 return Ok(None);
117 };
118 let n = match value {
119 ferrox_gguf::GgufValue::U8(v) => i64::from(*v),
120 ferrox_gguf::GgufValue::I8(v) => i64::from(*v),
121 ferrox_gguf::GgufValue::U16(v) => i64::from(*v),
122 ferrox_gguf::GgufValue::I16(v) => i64::from(*v),
123 ferrox_gguf::GgufValue::U32(v) => i64::from(*v),
124 ferrox_gguf::GgufValue::I32(v) => i64::from(*v),
125 ferrox_gguf::GgufValue::U64(v) => *v as i64,
126 ferrox_gguf::GgufValue::I64(v) => *v,
127 other => {
128 return Err(PoolingError::NotAnInteger {
129 key,
130 value: format!("{other:?}"),
131 })
132 }
133 };
134 Self::from_wire(&key, n)
135 }
136}
137
138pub fn pool(hidden: &[f32], n_embd: usize, ty: PoolingType) -> Result<Vec<f32>, PoolingError> {
148 if n_embd == 0 || hidden.is_empty() {
149 return Err(PoolingError::EmptySequence);
150 }
151 if !hidden.len().is_multiple_of(n_embd) {
152 return Err(PoolingError::RaggedHiddenStates {
153 len: hidden.len(),
154 n_embd,
155 });
156 }
157 let n_tokens = hidden.len() / n_embd;
158 Ok(match ty {
159 PoolingType::None => hidden.to_vec(),
160 PoolingType::Cls => hidden[..n_embd].to_vec(),
161 PoolingType::Last => hidden[(n_tokens - 1) * n_embd..].to_vec(),
162 PoolingType::Mean => {
163 let mut out = vec![0.0f32; n_embd];
164 for row in hidden.chunks_exact(n_embd) {
165 for (o, v) in out.iter_mut().zip(row) {
166 *o += *v;
167 }
168 }
169 let inv = 1.0 / n_tokens as f32;
170 for o in out.iter_mut() {
171 *o *= inv;
172 }
173 out
174 }
175 PoolingType::Rank => return Err(PoolingError::Unimplemented),
176 })
177}
178
179pub fn l2_normalize(v: &mut [f32]) {
185 let sum: f32 = v.iter().map(|x| x * x).sum();
186 if sum <= 0.0 {
187 return;
188 }
189 let inv = 1.0 / sum.sqrt();
190 for x in v.iter_mut() {
191 *x *= inv;
192 }
193}
194
195#[cfg(test)]
196mod tests {
197 use super::*;
198
199 #[test]
200 fn cls_takes_the_first_row_and_last_takes_the_last() {
201 let hidden = vec![1.0, 2.0, 10.0, 20.0, 100.0, 200.0];
202 assert_eq!(pool(&hidden, 2, PoolingType::Cls).unwrap(), vec![1.0, 2.0]);
203 assert_eq!(
204 pool(&hidden, 2, PoolingType::Last).unwrap(),
205 vec![100.0, 200.0]
206 );
207 assert_eq!(
208 pool(&hidden, 2, PoolingType::Mean).unwrap(),
209 vec![37.0, 74.0]
210 );
211 assert_eq!(pool(&hidden, 2, PoolingType::None).unwrap(), hidden);
212 }
213
214 #[test]
217 fn rank_refuses_by_name() {
218 let err = pool(&[1.0, 2.0], 2, PoolingType::Rank).unwrap_err();
219 assert!(matches!(err, PoolingError::Unimplemented));
220 assert!(err.to_string().contains("RANK"));
221 }
222
223 #[test]
224 fn empty_and_ragged_inputs_refuse() {
225 assert!(matches!(
226 pool(&[], 4, PoolingType::Cls),
227 Err(PoolingError::EmptySequence)
228 ));
229 assert!(matches!(
230 pool(&[1.0, 2.0, 3.0], 2, PoolingType::Cls),
231 Err(PoolingError::RaggedHiddenStates { len: 3, n_embd: 2 })
232 ));
233 }
234
235 #[test]
236 fn wire_values_match_llama_pooling_type() {
237 let k = "bert.pooling_type";
238 assert_eq!(PoolingType::from_wire(k, -1).unwrap(), None);
239 assert_eq!(
240 PoolingType::from_wire(k, 0).unwrap(),
241 Some(PoolingType::None)
242 );
243 assert_eq!(
244 PoolingType::from_wire(k, 1).unwrap(),
245 Some(PoolingType::Mean)
246 );
247 assert_eq!(
248 PoolingType::from_wire(k, 2).unwrap(),
249 Some(PoolingType::Cls)
250 );
251 assert_eq!(
252 PoolingType::from_wire(k, 3).unwrap(),
253 Some(PoolingType::Last)
254 );
255 assert_eq!(
256 PoolingType::from_wire(k, 4).unwrap(),
257 Some(PoolingType::Rank)
258 );
259 assert!(PoolingType::from_wire(k, 5).is_err());
260 }
261
262 #[test]
263 fn l2_normalize_makes_a_unit_vector_and_leaves_zero_alone() {
264 let mut v = vec![3.0f32, 4.0];
265 l2_normalize(&mut v);
266 assert!((v[0] - 0.6).abs() < 1e-6 && (v[1] - 0.8).abs() < 1e-6);
267 let mut z = vec![0.0f32; 3];
268 l2_normalize(&mut z);
269 assert_eq!(z, vec![0.0, 0.0, 0.0]);
270 }
271}