llm_kernel/embedding/
fastembed.rs1use std::path::PathBuf;
16use std::sync::Mutex;
17
18use crate::embedding::catalog::EmbeddingModel;
19use crate::embedding::types::{EmbeddingProvider, EmbeddingResult};
20use crate::error::{KernelError, Result};
21
22pub struct FastembedProvider {
28 inner: Mutex<fastembed::TextEmbedding>,
29 model: EmbeddingModel,
30}
31
32impl FastembedProvider {
33 pub fn new(model: EmbeddingModel, cache_dir: Option<PathBuf>) -> Result<Self> {
38 let mut options = fastembed::TextInitOptions::new(model.as_fastembed())
39 .with_show_download_progress(false);
40 if let Some(dir) = cache_dir {
41 options = options.with_cache_dir(dir);
42 }
43 let te = fastembed::TextEmbedding::try_new(options).map_err(KernelError::embedding)?;
44 Ok(Self {
45 inner: Mutex::new(te),
46 model,
47 })
48 }
49
50 #[cfg(all(feature = "embedding-fastembed-directml", target_os = "windows"))]
61 pub fn new_with_directml(model: EmbeddingModel, cache_dir: Option<PathBuf>) -> Result<Self> {
62 use ort::execution_providers::DirectMLExecutionProvider;
63 let mut options = fastembed::TextInitOptions::new(model.as_fastembed())
64 .with_show_download_progress(false)
65 .with_execution_providers(vec![DirectMLExecutionProvider::default().build()]);
66 if let Some(dir) = cache_dir {
67 options = options.with_cache_dir(dir);
68 }
69 let te = fastembed::TextEmbedding::try_new(options).map_err(KernelError::embedding)?;
70 Ok(Self {
71 inner: Mutex::new(te),
72 model,
73 })
74 }
75
76 #[cfg(all(feature = "embedding-fastembed-coreml", target_os = "macos"))]
82 pub fn new_with_coreml(model: EmbeddingModel, cache_dir: Option<PathBuf>) -> Result<Self> {
83 use ort::execution_providers::CoreMLExecutionProvider;
84 let mut options = fastembed::TextInitOptions::new(model.as_fastembed())
85 .with_show_download_progress(false)
86 .with_execution_providers(vec![CoreMLExecutionProvider::default().build()]);
87 if let Some(dir) = cache_dir {
88 options = options.with_cache_dir(dir);
89 }
90 let te = fastembed::TextEmbedding::try_new(options).map_err(KernelError::embedding)?;
91 Ok(Self {
92 inner: Mutex::new(te),
93 model,
94 })
95 }
96
97 pub fn with_max_length(
99 model: EmbeddingModel,
100 cache_dir: Option<PathBuf>,
101 max_length: usize,
102 ) -> Result<Self> {
103 let mut options = fastembed::TextInitOptions::new(model.as_fastembed())
104 .with_show_download_progress(false)
105 .with_max_length(max_length);
106 if let Some(dir) = cache_dir {
107 options = options.with_cache_dir(dir);
108 }
109 let te = fastembed::TextEmbedding::try_new(options).map_err(KernelError::embedding)?;
110 Ok(Self {
111 inner: Mutex::new(te),
112 model,
113 })
114 }
115}
116
117use super::types::text_preview;
118
119impl EmbeddingProvider for FastembedProvider {
120 fn dim(&self) -> usize {
121 self.model.dimension()
122 }
123
124 fn name(&self) -> &str {
125 self.model.as_str()
126 }
127
128 fn embed(&self, text: &str) -> Result<EmbeddingResult> {
129 self.embed_with_prefix(text, self.model.query_prefix())
130 }
131
132 fn embed_document(&self, text: &str) -> Result<EmbeddingResult> {
133 self.embed_with_prefix(text, self.model.doc_prefix())
136 }
137
138 fn embed_batch(&self, texts: &[&str]) -> Result<Vec<EmbeddingResult>> {
139 self.embed_batch_with_prefix(texts, self.model.query_prefix())
140 }
141
142 fn embed_documents(&self, texts: &[&str]) -> Result<Vec<EmbeddingResult>> {
143 self.embed_batch_with_prefix(texts, self.model.doc_prefix())
144 }
145}
146
147impl FastembedProvider {
148 fn embed_with_prefix(&self, text: &str, prefix: Option<&str>) -> Result<EmbeddingResult> {
150 let owned = match prefix {
151 Some(p) => format!("{p}{text}"),
152 None => text.to_string(),
153 };
154 let mut te = self
155 .inner
156 .lock()
157 .map_err(|e| KernelError::Embedding(format!("lock: {e}")))?;
158 let embeddings = te
159 .embed(vec![owned], None)
160 .map_err(KernelError::embedding)?;
161 let vector = embeddings
162 .into_iter()
163 .next()
164 .ok_or_else(|| KernelError::Embedding("empty embedding output".into()))?;
165
166 Ok(EmbeddingResult {
167 vector,
168 text_preview: text_preview(text),
169 })
170 }
171
172 fn embed_batch_with_prefix(
174 &self,
175 texts: &[&str],
176 prefix: Option<&str>,
177 ) -> Result<Vec<EmbeddingResult>> {
178 if texts.is_empty() {
179 return Ok(vec![]);
180 }
181 let prepared: Vec<String> = texts
182 .iter()
183 .map(|t| match prefix {
184 Some(p) => format!("{p}{t}"),
185 None => t.to_string(),
186 })
187 .collect();
188
189 let mut te = self
190 .inner
191 .lock()
192 .map_err(|e| KernelError::Embedding(format!("lock: {e}")))?;
193 let embeddings = te.embed(prepared, None).map_err(KernelError::embedding)?;
194
195 Ok(embeddings
196 .into_iter()
197 .zip(texts.iter())
198 .map(|(vector, &text)| EmbeddingResult {
199 vector,
200 text_preview: text_preview(text),
201 })
202 .collect())
203 }
204}
205
206#[cfg(test)]
207mod tests {
208 use super::*;
209
210 #[test]
211 fn provider_name_matches_model() {
212 for &m in EmbeddingModel::ALL {
215 let fe = m.as_fastembed();
217 assert_eq!(format!("{fe:?}"), m.as_str());
218 }
219 }
220
221 #[test]
222 #[ignore = "requires model download"]
223 fn embed_single_text() {
224 let dir = tempfile::tempdir().unwrap();
225 let provider = FastembedProvider::new(
226 EmbeddingModel::BGESmallENV15,
227 Some(dir.path().to_path_buf()),
228 )
229 .unwrap();
230 let result = provider.embed("hello world").unwrap();
231 assert_eq!(result.vector.len(), 384);
232 assert!(!result.vector.is_empty());
233 }
234
235 #[test]
236 #[ignore = "requires model download"]
237 fn embed_batch_texts() {
238 let dir = tempfile::tempdir().unwrap();
239 let provider = FastembedProvider::new(
240 EmbeddingModel::BGESmallENV15,
241 Some(dir.path().to_path_buf()),
242 )
243 .unwrap();
244 let results = provider
245 .embed_batch(&["hello", "world", "foo bar"])
246 .unwrap();
247 assert_eq!(results.len(), 3);
248 for r in &results {
249 assert_eq!(r.vector.len(), 384);
250 }
251 }
252}