av_loopdetect/
onnx_embed.rs1use crate::embed::Embedder;
11use std::path::Path;
12use std::sync::Arc;
13use tract_onnx::prelude::*;
14
15type RunnableOnnxModel = Arc<TypedRunnableModel>;
16
17pub struct OnnxEmbedder {
19 model: RunnableOnnxModel,
20 tokenizer: tokenizers::Tokenizer,
21 input_count: usize,
22 dim: usize,
23}
24
25impl OnnxEmbedder {
26 pub fn load(path: &Path, tokenizer_path: &Path, dim: usize) -> Result<Self, String> {
28 if dim == 0 {
29 return Err("ONNX embedding dimension must be greater than zero".to_owned());
30 }
31 let model = tract_onnx::onnx()
32 .model_for_path(path)
33 .map_err(|error| error.to_string())?
34 .into_optimized()
35 .map_err(|error| error.to_string())?;
36 let input_count = model.input_outlets().map_err(|error| error.to_string())?.len();
37 if !(2..=3).contains(&input_count) {
38 return Err(format!(
39 "ONNX sentence model must have 2 or 3 inputs, found {input_count}"
40 ));
41 }
42 let model = model.into_runnable().map_err(|error| error.to_string())?;
43 let tokenizer =
44 tokenizers::Tokenizer::from_file(tokenizer_path).map_err(|error| error.to_string())?;
45 Ok(Self {
46 model,
47 tokenizer,
48 input_count,
49 dim,
50 })
51 }
52
53 fn infer(&self, text: &str) -> Result<Vec<f32>, String> {
54 let encoding = self
55 .tokenizer
56 .encode(text, true)
57 .map_err(|error| error.to_string())?;
58 let ids: Vec<i64> = encoding
59 .get_ids()
60 .iter()
61 .take(512)
62 .map(|id| i64::from(*id))
63 .collect();
64 if ids.is_empty() {
65 return Ok(vec![0.0; self.dim]);
66 }
67 let len = ids.len();
68 let input =
69 tract_ndarray::Array2::from_shape_vec((1, len), ids).map_err(|error| error.to_string())?;
70 let mask_values: Vec<i64> = encoding
71 .get_attention_mask()
72 .iter()
73 .take(len)
74 .map(|value| i64::from(*value))
75 .collect();
76 let mask = tract_ndarray::Array2::from_shape_vec((1, len), mask_values.clone())
77 .map_err(|error| error.to_string())?;
78 let inputs = if self.input_count == 3 {
79 let token_types: Vec<i64> = encoding
80 .get_type_ids()
81 .iter()
82 .take(len)
83 .map(|value| i64::from(*value))
84 .collect();
85 let token_types = tract_ndarray::Array2::from_shape_vec((1, len), token_types)
86 .map_err(|error| error.to_string())?;
87 tvec!(
88 Tensor::from(input).into(),
89 Tensor::from(mask).into(),
90 Tensor::from(token_types).into()
91 )
92 } else {
93 tvec!(Tensor::from(input).into(), Tensor::from(mask).into())
94 };
95 let outputs = self.model.run(inputs).map_err(|error| error.to_string())?;
96 let output = outputs
97 .first()
98 .ok_or_else(|| "ONNX model returned no output".to_owned())?;
99 let view = output
100 .to_plain_array_view::<f32>()
101 .map_err(|error| error.to_string())?;
102 let mut vector = pool_output(view, &mask_values, self.dim)?;
103 let norm: f32 = vector.iter().map(|value| value * value).sum::<f32>().sqrt();
104 if norm > 0.0 {
105 for value in &mut vector {
106 *value /= norm;
107 }
108 }
109 Ok(vector)
110 }
111}
112
113fn pool_output(
114 output: tract_ndarray::ArrayViewD<'_, f32>,
115 attention_mask: &[i64],
116 expected_dim: usize,
117) -> Result<Vec<f32>, String> {
118 let shape = output.shape();
119 let vector = match shape {
120 [width] if *width == expected_dim => output.iter().copied().collect(),
121 [1, width] if *width == expected_dim => output.iter().copied().collect(),
122 [1, tokens, width] if *width == expected_dim && *tokens == attention_mask.len() => {
123 let mut pooled = vec![0.0f32; expected_dim];
124 let mut weight = 0.0f32;
125 for (token, mask) in attention_mask.iter().enumerate() {
126 if *mask <= 0 {
127 continue;
128 }
129 let token_weight = *mask as f32;
130 weight += token_weight;
131 for (feature, value) in pooled.iter_mut().enumerate() {
132 let index = tract_ndarray::IxDyn(&[0, token, feature]);
133 let embedding = output
134 .get(index)
135 .ok_or_else(|| "ONNX output index escaped validated shape".to_owned())?;
136 *value += *embedding * token_weight;
137 }
138 }
139 if weight == 0.0 {
140 return Err("ONNX attention mask contains no active tokens".to_owned());
141 }
142 for value in &mut pooled {
143 *value /= weight;
144 }
145 pooled
146 }
147 _ => {
148 return Err(format!(
149 "ONNX output shape {shape:?} is incompatible with embedding dimension {expected_dim} and token count {}",
150 attention_mask.len()
151 ));
152 }
153 };
154 if vector.iter().any(|value| !value.is_finite()) {
155 return Err("ONNX output contains non-finite values".to_owned());
156 }
157 Ok(vector)
158}
159
160impl Embedder for OnnxEmbedder {
161 fn dim(&self) -> usize {
162 self.dim
163 }
164
165 fn embed(&self, text: &str) -> Vec<f32> {
166 self.infer(text).unwrap_or_else(|error| {
167 tracing::warn!(%error, dim = self.dim, "ONNX inference failed; returning zero vector");
168 vec![0.0; self.dim]
169 })
170 }
171
172 fn try_embed(&self, text: &str) -> Result<Vec<f32>, String> {
173 self.infer(text)
174 }
175}
176
177#[cfg(test)]
178mod tests {
179 #![allow(clippy::indexing_slicing, clippy::unwrap_used)]
180
181 use super::*;
182
183 #[test]
184 fn token_embeddings_use_masked_mean_pooling() {
185 let output = tract_ndarray::Array3::from_shape_vec((1, 3, 2), vec![1.0, 0.0, 1.0, 2.0, 100.0, 100.0])
186 .unwrap()
187 .into_dyn();
188 let mut vector = pool_output(output.view(), &[1, 1, 0], 2).unwrap();
189 let norm = vector.iter().map(|value| value * value).sum::<f32>().sqrt();
190 for value in &mut vector {
191 *value /= norm;
192 }
193 let expected = 1.0 / 2.0f32.sqrt();
194 assert!((vector[0] - expected).abs() < 1e-6);
195 assert!((vector[1] - expected).abs() < 1e-6);
196 }
197
198 #[test]
199 fn pooled_export_requires_exact_embedding_width() {
200 let output = tract_ndarray::Array2::from_shape_vec((1, 3), vec![1.0, 2.0, 3.0])
201 .unwrap()
202 .into_dyn();
203 let error = pool_output(output.view(), &[1], 2).unwrap_err();
204 assert!(error.contains("output shape"));
205 }
206
207 #[test]
208 fn token_output_requires_matching_attention_mask() {
209 let output = tract_ndarray::Array3::zeros((1, 2, 3)).into_dyn();
210 let error = pool_output(output.view(), &[1], 3).unwrap_err();
211 assert!(error.contains("token count 1"));
212 }
213}