Skip to main content

ragfs_embed/
candle.rs

1//! GTE-small embedder using Candle.
2//!
3//! Uses thenlper/gte-small model for text embeddings:
4//! - 384 dimensions
5//! - 512 max tokens
6//! - BERT architecture
7
8use 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
20/// Model identifier on `HuggingFace` Hub.
21const MODEL_ID: &str = "thenlper/gte-small";
22
23/// Resolve a user-facing model name to the implemented Hugging Face id.
24///
25/// RAGFS currently implements only `thenlper/gte-small` (alias: `gte-small`).
26pub 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
35/// Embedding dimension for gte-small.
36const EMBEDDING_DIM: usize = 384;
37
38/// Maximum sequence length.
39const MAX_TOKENS: usize = 512;
40
41/// GTE-small embedder using Candle.
42pub struct CandleEmbedder {
43    /// Device to run inference on (CPU or CUDA)
44    device: Device,
45    /// Loaded model
46    model: Arc<RwLock<Option<BertModel>>>,
47    /// Tokenizer
48    tokenizer: Arc<RwLock<Option<Tokenizer>>>,
49    /// Model configuration
50    config: Arc<RwLock<Option<Config>>>,
51    /// Cache directory for models
52    cache_dir: PathBuf,
53    /// Whether model is initialized
54    initialized: Arc<RwLock<bool>>,
55}
56
57impl CandleEmbedder {
58    /// Create a new `CandleEmbedder` with the default gte-small model.
59    ///
60    /// GPU is used when available. Prefer [`Self::try_new`] to honor config.
61    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    /// Create an embedder from config (`model`, `use_gpu`).
67    ///
68    /// Unsupported models fail immediately with a clear error — before download.
69    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    /// Create with specific device.
89    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    /// Whether inference will run on CPU (config `use_gpu = false`, or no GPU).
101    #[must_use]
102    pub fn device_is_cpu(&self) -> bool {
103        self.device.is_cpu()
104    }
105
106    /// Initialize the model (download if needed, load into memory).
107    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        // Download model files from HuggingFace Hub into the configured cache dir
118        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        // Download tokenizer
126        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        // Download model config
133        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        // Download model weights
140        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        // Load tokenizer
147        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        // Load config
152        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        // Load model weights
159        debug!("Loading model weights...");
160        // SAFETY: The safetensors file is downloaded from HuggingFace Hub and is trusted.
161        // Memory mapping is safe for read-only access to model weights.
162        #[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        // Store in instance
172        {
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    /// Mean pooling with attention mask.
194    fn mean_pooling(
195        &self,
196        token_embeddings: &Tensor,
197        attention_mask: &Tensor,
198    ) -> Result<Tensor, EmbedError> {
199        // Expand attention mask to match embedding dimensions
200        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        // Masked sum
209        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        // Count non-masked tokens
218        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        // Mean
225        let mean = sum
226            .div(&mask_sum)
227            .map_err(|e| EmbedError::Inference(format!("div failed: {e}")))?;
228
229        Ok(mean)
230    }
231
232    /// L2 normalize embeddings.
233    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    /// Encode a batch of texts.
252    async fn encode_batch(
253        &self,
254        texts: &[&str],
255        normalize: bool,
256    ) -> Result<Vec<EmbeddingOutput>, EmbedError> {
257        // Ensure initialized
258        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        // Tokenize all texts
271        let encodings = tokenizer
272            .encode_batch(texts.to_vec(), true)
273            .map_err(|e| EmbedError::Inference(format!("Tokenization failed: {e}")))?;
274
275        // Find max length for padding
276        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        // Prepare input tensors
284        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            // Add IDs with padding
295            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); // PAD token
302                    attention_mask_vec.push(0);
303                    token_type_ids_vec.push(0);
304                }
305            }
306        }
307
308        let batch_size = texts.len();
309
310        // Create tensors
311        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        // Run model
327        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        // Mean pooling
332        let pooled = self.mean_pooling(&output, &attention_mask)?;
333
334        // Normalize if requested
335        let final_embeddings = if normalize {
336            self.normalize(&pooled)?
337        } else {
338            pooled
339        };
340
341        // Convert to Vec<EmbeddingOutput>
342        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        // Process in batches
395        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        // For GTE models, queries and documents use the same embedding process
411        // Some models use different prefixes, but GTE doesn't need that
412        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] // Requires model download
469    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        // Check normalization (should have unit length)
487        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}