Skip to main content

av_loopdetect/
onnx_embed.rs

1//! ONNX embedder via tract (pure Rust, no PyTorch/Python runtime — brief §8).
2//!
3//! Loads a MiniLM-class sentence-embedding ONNX model from a configured path.
4//! Deployment note: model files are customer-supplied artifacts (air-gapped
5//! installs cannot download); the default deployment uses [`crate::HashEmbedder`]
6//! and swapping to ONNX is a config change, not a code change (plan D6).
7//!
8//! Tokenization is supplied by the model's Hugging Face `tokenizer.json`.
9
10use crate::embed::Embedder;
11use std::path::Path;
12use std::sync::Arc;
13use tract_onnx::prelude::*;
14
15type RunnableOnnxModel = Arc<TypedRunnableModel>;
16
17/// Embedder backed by an ONNX sentence-embedding model.
18pub struct OnnxEmbedder {
19    model: RunnableOnnxModel,
20    tokenizer: tokenizers::Tokenizer,
21    input_count: usize,
22    dim: usize,
23}
24
25impl OnnxEmbedder {
26    /// Load a model and its paired tokenizer. `dim` must match the output width.
27    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}