1use async_trait::async_trait;
9use candle_core::{DType, Device, Tensor};
10use candle_nn::VarBuilder;
11use candle_transformers::models::bert::{BertModel, Config};
12use hf_hub::{Repo, RepoType};
13use ragfs_core::{EmbedError, Embedder, EmbeddingConfig, EmbeddingOutput, Modality};
14use std::path::PathBuf;
15use std::sync::Arc;
16use tokenizers::Tokenizer;
17use tokio::sync::RwLock;
18use tracing::{debug, info};
19
20const MODEL_ID: &str = "thenlper/gte-small";
22
23pub fn resolve_supported_model(model: &str) -> Result<&'static str, EmbedError> {
27 match model.trim() {
28 "thenlper/gte-small" | "gte-small" => Ok(MODEL_ID),
29 other => Err(EmbedError::ModelLoad(format!(
30 "Unsupported embedding model '{other}'. RAGFS currently supports only 'thenlper/gte-small' (alias: 'gte-small')."
31 ))),
32 }
33}
34
35const EMBEDDING_DIM: usize = 384;
37
38const MAX_TOKENS: usize = 512;
40
41pub struct CandleEmbedder {
43 device: Device,
45 model: Arc<RwLock<Option<BertModel>>>,
47 tokenizer: Arc<RwLock<Option<Tokenizer>>>,
49 config: Arc<RwLock<Option<Config>>>,
51 cache_dir: PathBuf,
53 initialized: Arc<RwLock<bool>>,
55}
56
57impl CandleEmbedder {
58 pub fn new(cache_dir: PathBuf) -> Self {
62 Self::try_new(cache_dir, MODEL_ID, true)
63 .expect("default model thenlper/gte-small is supported")
64 }
65
66 pub fn try_new(cache_dir: PathBuf, model: &str, use_gpu: bool) -> Result<Self, EmbedError> {
70 let _model_id = resolve_supported_model(model)?;
71 let device = if use_gpu {
72 Device::cuda_if_available(0).unwrap_or(Device::Cpu)
73 } else {
74 Device::Cpu
75 };
76 info!("CandleEmbedder using device: {:?}", device);
77
78 Ok(Self {
79 device,
80 model: Arc::new(RwLock::new(None)),
81 tokenizer: Arc::new(RwLock::new(None)),
82 config: Arc::new(RwLock::new(None)),
83 cache_dir,
84 initialized: Arc::new(RwLock::new(false)),
85 })
86 }
87
88 pub fn with_device(cache_dir: PathBuf, device: Device) -> Self {
90 Self {
91 device,
92 model: Arc::new(RwLock::new(None)),
93 tokenizer: Arc::new(RwLock::new(None)),
94 config: Arc::new(RwLock::new(None)),
95 cache_dir,
96 initialized: Arc::new(RwLock::new(false)),
97 }
98 }
99
100 #[must_use]
102 pub fn device_is_cpu(&self) -> bool {
103 self.device.is_cpu()
104 }
105
106 pub async fn init(&self) -> Result<(), EmbedError> {
108 {
109 let initialized = self.initialized.read().await;
110 if *initialized {
111 return Ok(());
112 }
113 }
114
115 info!("Initializing CandleEmbedder with model: {}", MODEL_ID);
116
117 let api = hf_hub::api::tokio::ApiBuilder::new()
119 .with_cache_dir(self.cache_dir.clone())
120 .build()
121 .map_err(|e| EmbedError::ModelLoad(format!("Failed to create HF API: {e}")))?;
122
123 let repo = api.repo(Repo::new(MODEL_ID.to_string(), RepoType::Model));
124
125 debug!("Downloading tokenizer...");
127 let tokenizer_path = repo
128 .get("tokenizer.json")
129 .await
130 .map_err(|e| EmbedError::ModelLoad(format!("Failed to download tokenizer: {e}")))?;
131
132 debug!("Downloading config...");
134 let config_path = repo
135 .get("config.json")
136 .await
137 .map_err(|e| EmbedError::ModelLoad(format!("Failed to download config: {e}")))?;
138
139 debug!("Downloading model weights...");
141 let weights_path = repo
142 .get("model.safetensors")
143 .await
144 .map_err(|e| EmbedError::ModelLoad(format!("Failed to download weights: {e}")))?;
145
146 debug!("Loading tokenizer...");
148 let tokenizer = Tokenizer::from_file(&tokenizer_path)
149 .map_err(|e| EmbedError::ModelLoad(format!("Failed to load tokenizer: {e}")))?;
150
151 debug!("Loading config...");
153 let config_str = std::fs::read_to_string(&config_path)
154 .map_err(|e| EmbedError::ModelLoad(format!("Failed to read config: {e}")))?;
155 let config: Config = serde_json::from_str(&config_str)
156 .map_err(|e| EmbedError::ModelLoad(format!("Failed to parse config: {e}")))?;
157
158 debug!("Loading model weights...");
160 #[allow(unsafe_code)]
163 let vb = unsafe {
164 VarBuilder::from_mmaped_safetensors(&[weights_path], DType::F32, &self.device)
165 .map_err(|e| EmbedError::ModelLoad(format!("Failed to load weights: {e}")))?
166 };
167
168 let model = BertModel::load(vb, &config)
169 .map_err(|e| EmbedError::ModelLoad(format!("Failed to create BERT model: {e}")))?;
170
171 {
173 let mut tok = self.tokenizer.write().await;
174 *tok = Some(tokenizer);
175 }
176 {
177 let mut cfg = self.config.write().await;
178 *cfg = Some(config);
179 }
180 {
181 let mut mdl = self.model.write().await;
182 *mdl = Some(model);
183 }
184 {
185 let mut init = self.initialized.write().await;
186 *init = true;
187 }
188
189 info!("CandleEmbedder initialized successfully");
190 Ok(())
191 }
192
193 fn mean_pooling(
195 &self,
196 token_embeddings: &Tensor,
197 attention_mask: &Tensor,
198 ) -> Result<Tensor, EmbedError> {
199 let mask = attention_mask
201 .unsqueeze(2)
202 .map_err(|e| EmbedError::Inference(format!("unsqueeze failed: {e}")))?
203 .broadcast_as(token_embeddings.shape())
204 .map_err(|e| EmbedError::Inference(format!("broadcast failed: {e}")))?
205 .to_dtype(DType::F32)
206 .map_err(|e| EmbedError::Inference(format!("dtype conversion failed: {e}")))?;
207
208 let masked = token_embeddings
210 .mul(&mask)
211 .map_err(|e| EmbedError::Inference(format!("mul failed: {e}")))?;
212
213 let sum = masked
214 .sum(1)
215 .map_err(|e| EmbedError::Inference(format!("sum failed: {e}")))?;
216
217 let mask_sum = mask
219 .sum(1)
220 .map_err(|e| EmbedError::Inference(format!("mask sum failed: {e}")))?
221 .clamp(1e-9, f64::MAX)
222 .map_err(|e| EmbedError::Inference(format!("clamp failed: {e}")))?;
223
224 let mean = sum
226 .div(&mask_sum)
227 .map_err(|e| EmbedError::Inference(format!("div failed: {e}")))?;
228
229 Ok(mean)
230 }
231
232 fn normalize(&self, embeddings: &Tensor) -> Result<Tensor, EmbedError> {
234 let norm = embeddings
235 .sqr()
236 .map_err(|e| EmbedError::Inference(format!("sqr failed: {e}")))?
237 .sum_keepdim(1)
238 .map_err(|e| EmbedError::Inference(format!("sum_keepdim failed: {e}")))?
239 .sqrt()
240 .map_err(|e| EmbedError::Inference(format!("sqrt failed: {e}")))?
241 .clamp(1e-12, f64::MAX)
242 .map_err(|e| EmbedError::Inference(format!("clamp failed: {e}")))?;
243
244 let normalized = embeddings
245 .broadcast_div(&norm)
246 .map_err(|e| EmbedError::Inference(format!("div failed: {e}")))?;
247
248 Ok(normalized)
249 }
250
251 async fn encode_batch(
253 &self,
254 texts: &[&str],
255 normalize: bool,
256 ) -> Result<Vec<EmbeddingOutput>, EmbedError> {
257 self.init().await?;
259
260 let tokenizer = self.tokenizer.read().await;
261 let tokenizer = tokenizer
262 .as_ref()
263 .ok_or_else(|| EmbedError::Inference("Tokenizer not loaded".to_string()))?;
264
265 let model = self.model.read().await;
266 let model = model
267 .as_ref()
268 .ok_or_else(|| EmbedError::Inference("Model not loaded".to_string()))?;
269
270 let encodings = tokenizer
272 .encode_batch(texts.to_vec(), true)
273 .map_err(|e| EmbedError::Inference(format!("Tokenization failed: {e}")))?;
274
275 let max_len = encodings
277 .iter()
278 .map(tokenizers::Encoding::len)
279 .max()
280 .unwrap_or(0);
281 let max_len = max_len.min(MAX_TOKENS);
282
283 let mut input_ids_vec: Vec<u32> = Vec::new();
285 let mut attention_mask_vec: Vec<u32> = Vec::new();
286 let mut token_type_ids_vec: Vec<u32> = Vec::new();
287 let mut token_counts = Vec::new();
288
289 for encoding in &encodings {
290 let ids = encoding.get_ids();
291 let len = ids.len().min(max_len);
292 token_counts.push(len);
293
294 for i in 0..max_len {
296 if i < len {
297 input_ids_vec.push(ids[i]);
298 attention_mask_vec.push(1);
299 token_type_ids_vec.push(0);
300 } else {
301 input_ids_vec.push(0); attention_mask_vec.push(0);
303 token_type_ids_vec.push(0);
304 }
305 }
306 }
307
308 let batch_size = texts.len();
309
310 let input_ids = Tensor::from_vec(input_ids_vec, (batch_size, max_len), &self.device)
312 .map_err(|e| {
313 EmbedError::Inference(format!("Failed to create input_ids tensor: {e}"))
314 })?;
315
316 let attention_mask =
317 Tensor::from_vec(attention_mask_vec, (batch_size, max_len), &self.device).map_err(
318 |e| EmbedError::Inference(format!("Failed to create attention_mask tensor: {e}")),
319 )?;
320
321 let token_type_ids =
322 Tensor::from_vec(token_type_ids_vec, (batch_size, max_len), &self.device).map_err(
323 |e| EmbedError::Inference(format!("Failed to create token_type_ids tensor: {e}")),
324 )?;
325
326 let output = model
328 .forward(&input_ids, &token_type_ids, Some(&attention_mask))
329 .map_err(|e| EmbedError::Inference(format!("Model forward failed: {e}")))?;
330
331 let pooled = self.mean_pooling(&output, &attention_mask)?;
333
334 let final_embeddings = if normalize {
336 self.normalize(&pooled)?
337 } else {
338 pooled
339 };
340
341 let mut results = Vec::with_capacity(batch_size);
343
344 for i in 0..batch_size {
345 let embedding = final_embeddings
346 .get(i)
347 .map_err(|e| EmbedError::Inference(format!("Failed to get embedding {i}: {e}")))?
348 .to_vec1::<f32>()
349 .map_err(|e| EmbedError::Inference(format!("Failed to convert to vec: {e}")))?;
350
351 results.push(EmbeddingOutput {
352 embedding,
353 token_count: token_counts[i],
354 });
355 }
356
357 Ok(results)
358 }
359}
360
361#[async_trait]
362impl Embedder for CandleEmbedder {
363 fn model_name(&self) -> &str {
364 MODEL_ID
365 }
366
367 fn dimension(&self) -> usize {
368 EMBEDDING_DIM
369 }
370
371 fn max_tokens(&self) -> usize {
372 MAX_TOKENS
373 }
374
375 fn modalities(&self) -> &[Modality] {
376 &[Modality::Text]
377 }
378
379 async fn embed_text(
380 &self,
381 texts: &[&str],
382 config: &EmbeddingConfig,
383 ) -> Result<Vec<EmbeddingOutput>, EmbedError> {
384 if texts.is_empty() {
385 return Ok(Vec::new());
386 }
387
388 debug!(
389 "Embedding {} texts with batch_size {}",
390 texts.len(),
391 config.batch_size
392 );
393
394 let mut all_results = Vec::with_capacity(texts.len());
396
397 for chunk in texts.chunks(config.batch_size) {
398 let batch_results = self.encode_batch(chunk, config.normalize).await?;
399 all_results.extend(batch_results);
400 }
401
402 Ok(all_results)
403 }
404
405 async fn embed_query(
406 &self,
407 query: &str,
408 config: &EmbeddingConfig,
409 ) -> Result<EmbeddingOutput, EmbedError> {
410 let results = self.embed_text(&[query], config).await?;
413 results
414 .into_iter()
415 .next()
416 .ok_or_else(|| EmbedError::Inference("Empty embedding result".to_string()))
417 }
418}
419
420#[cfg(test)]
421mod tests {
422 use super::*;
423 use tempfile::tempdir;
424
425 #[test]
426 fn test_unsupported_model_errors_before_download() {
427 let cache_dir = tempdir().unwrap();
428 let result =
429 CandleEmbedder::try_new(cache_dir.path().to_path_buf(), "jina-embeddings-v3", false);
430 let err = match result {
431 Ok(_) => panic!("unsupported model should not construct an embedder"),
432 Err(e) => e,
433 };
434 let message = err.to_string();
435 assert!(
436 message.contains("jina-embeddings-v3"),
437 "error should name the requested model: {message}"
438 );
439 assert!(
440 message.contains("thenlper/gte-small"),
441 "error should name the supported model: {message}"
442 );
443 }
444
445 #[test]
446 fn test_gte_small_alias_is_accepted() {
447 let cache_dir = tempdir().unwrap();
448 let embedder =
449 CandleEmbedder::try_new(cache_dir.path().to_path_buf(), "gte-small", false).unwrap();
450 assert_eq!(embedder.model_name(), "thenlper/gte-small");
451 assert!(
452 embedder.device_is_cpu(),
453 "use_gpu=false must select the CPU device"
454 );
455 }
456
457 #[test]
458 fn test_resolve_supported_model() {
459 assert_eq!(
460 resolve_supported_model("thenlper/gte-small").unwrap(),
461 MODEL_ID
462 );
463 assert_eq!(resolve_supported_model("gte-small").unwrap(), MODEL_ID);
464 assert!(resolve_supported_model("jina-embeddings-v3").is_err());
465 }
466
467 #[tokio::test]
468 #[ignore] async fn test_candle_embedder() {
470 let cache_dir = tempdir().unwrap();
471 let embedder = CandleEmbedder::new(cache_dir.path().to_path_buf());
472
473 embedder.init().await.unwrap();
474
475 assert_eq!(embedder.dimension(), 384);
476 assert_eq!(embedder.model_name(), "thenlper/gte-small");
477
478 let config = EmbeddingConfig::default();
479 let texts = &["Hello world", "This is a test"];
480
481 let results = embedder.embed_text(texts, &config).await.unwrap();
482 assert_eq!(results.len(), 2);
483 assert_eq!(results[0].embedding.len(), 384);
484 assert_eq!(results[1].embedding.len(), 384);
485
486 let norm: f32 = results[0]
488 .embedding
489 .iter()
490 .map(|x| x * x)
491 .sum::<f32>()
492 .sqrt();
493 assert!((norm - 1.0).abs() < 0.01);
494 }
495}