Skip to main content

lc_embeddings/
local.rs

1// lc-embeddings/src/local.rs
2//! Local embedding implementations
3//!
4//! Contains two implementations:
5//! - `BagOfWordsEmbeddings`: Lightweight word-frequency hash embedding (pure Rust, no external deps), always available
6//! - `LocalEmbeddings`: ONNX Runtime-based neural network embedding (requires `local-embeddings` feature)
7//!
8//! `BagOfWordsEmbeddings` is suitable for offline, privacy, zero-cost coarse-grained retrieval;
9//! `LocalEmbeddings` is suitable for high-quality semantic embedding scenarios (e.g., BGE/E5 models).
10
11use async_trait::async_trait;
12
13#[cfg(feature = "local-embeddings")]
14use std::path::Path;
15
16use crate::{EmbeddingError, Embeddings};
17
18// ---------------------------------------------------------------------------
19// BagOfWordsEmbeddings — word-frequency hash + L2 normalization (always available)
20// ---------------------------------------------------------------------------
21
22/// Lightweight local embedding (word-frequency hash + L2 normalization)
23///
24/// Based on word frequency + hashing, no API calls, suitable for offline, privacy, zero-cost coarse-grained retrieval.
25///
26/// Note: This is a lightweight implementation (bag-of-words hash) with limited semantic quality;
27/// for high-quality neural network embeddings (BGE/E5 via `ort`), enable the `local-embeddings` feature
28/// and use [`LocalEmbeddings`].
29pub struct BagOfWordsEmbeddings {
30    dim: usize,
31}
32
33impl BagOfWordsEmbeddings {
34    /// Create local embedding with specified dimension
35    pub fn new(dim: usize) -> Self {
36        Self { dim: dim.max(1) }
37    }
38
39    /// Default dimension 256
40    pub fn default_dim() -> Self {
41        Self::new(256)
42    }
43
44    /// Tokenize: English by non-alphanumeric split (lowercased), Chinese/non-ASCII by single character
45    fn tokenize(text: &str) -> Vec<String> {
46        let mut tokens = Vec::new();
47        let mut current = String::new();
48        for c in text.chars() {
49            if c.is_alphanumeric() {
50                if c.is_ascii() {
51                    current.push(c.to_ascii_lowercase());
52                } else {
53                    // Non-ASCII (Chinese etc.) single character as token
54                    if !current.is_empty() {
55                        tokens.push(std::mem::take(&mut current));
56                    }
57                    tokens.push(c.to_string());
58                }
59            } else if !current.is_empty() {
60                tokens.push(std::mem::take(&mut current));
61            }
62        }
63        if !current.is_empty() {
64            tokens.push(current);
65        }
66        tokens
67    }
68
69    /// FNV-1a hash
70    fn hash(s: &str) -> u64 {
71        let mut h: u64 = 0xcbf29ce484222325;
72        for b in s.bytes() {
73            h ^= b as u64;
74            h = h.wrapping_mul(0x100000001b3);
75        }
76        h
77    }
78
79    /// Compute embedding vector (word-frequency hash + L2 normalization)
80    fn embed(&self, text: &str) -> Vec<f32> {
81        let mut v = vec![0.0f32; self.dim];
82        for token in Self::tokenize(text) {
83            let idx = (Self::hash(&token) as usize) % self.dim;
84            v[idx] += 1.0;
85        }
86        // L2 normalization
87        let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
88        if norm > 0.0 {
89            for x in &mut v {
90                *x /= norm;
91            }
92        }
93        v
94    }
95}
96
97impl Default for BagOfWordsEmbeddings {
98    fn default() -> Self {
99        Self::default_dim()
100    }
101}
102
103#[async_trait]
104impl Embeddings for BagOfWordsEmbeddings {
105    async fn embed_query(&self, text: &str) -> Result<Vec<f32>, EmbeddingError> {
106        if text.trim().is_empty() {
107            return Err(EmbeddingError::EmptyInput);
108        }
109        Ok(self.embed(text))
110    }
111
112    fn dimension(&self) -> usize {
113        self.dim
114    }
115
116    fn model_name(&self) -> &str {
117        "local-bow"
118    }
119}
120
121// ---------------------------------------------------------------------------
122// LocalEmbeddings — ONNX Runtime neural network embedding (requires local-embeddings feature)
123// ---------------------------------------------------------------------------
124
125#[cfg(feature = "local-embeddings")]
126mod nn {
127    use super::*;
128    use ort::value::Tensor;
129    use std::sync::RwLock;
130
131    /// ONNX Runtime-based local neural network embedding
132    ///
133    /// Supports any ONNX format embedding model (e.g., BGE/E5), runs inference locally,
134    /// no external API calls needed, suitable for privacy-sensitive or offline scenarios.
135    ///
136    /// # Example
137    ///
138    /// ```ignore
139    /// use lc_embeddings::LocalEmbeddings;
140    ///
141    /// let embedder = LocalEmbeddings::from_file("model.onnx")?;
142    /// let vec = embedder.embed_query("hello world").await?;
143    /// ```
144    pub struct LocalEmbeddings {
145        // ort 2.0.0-rc.12's Session::run() requires &mut self,
146        // use RwLock to get mutable reference in &self methods while satisfying Send + Sync
147        session: RwLock<ort::session::Session>,
148        dim: usize,
149        model_name: String,
150    }
151
152    impl LocalEmbeddings {
153        /// Load from ONNX model file
154        ///
155        /// # Arguments
156        /// * `model_path` - Path to ONNX model file
157        ///
158        /// # Returns
159        /// `LocalEmbeddings` instance on success, `EmbeddingError` on failure
160        pub fn from_file(model_path: impl AsRef<Path>) -> Result<Self, EmbeddingError> {
161            let path = model_path.as_ref();
162            let model_name = path
163                .file_stem()
164                .and_then(|s| s.to_str())
165                .unwrap_or("unknown")
166                .to_string();
167
168            let session = ort::session::Session::builder()
169                .map_err(|e| {
170                    EmbeddingError::ApiError(format!("Failed to create ONNX SessionBuilder: {}", e))
171                })?
172                .commit_from_file(path)
173                .map_err(|e| {
174                    EmbeddingError::ApiError(format!(
175                        "Failed to load ONNX model ({}): {}",
176                        path.display(),
177                        e
178                    ))
179                })?;
180
181            // Infer dimension from model output
182            let dim = Self::infer_dimension(&session)?;
183
184            Ok(Self {
185                session: RwLock::new(session),
186                dim,
187                model_name,
188            })
189        }
190
191        /// Infer embedding dimension from ONNX session output info
192        fn infer_dimension(session: &ort::session::Session) -> Result<usize, EmbeddingError> {
193            let outputs = session.outputs();
194            if outputs.is_empty() {
195                return Err(EmbeddingError::ParseError(
196                    "ONNX model has no output nodes".to_string(),
197                ));
198            }
199
200            // Get first output's ValueType, infer dimension from shape
201            let dtype = outputs[0].dtype();
202            let shape = dtype.tensor_shape().ok_or_else(|| {
203                EmbeddingError::ParseError("Output is not a Tensor type".to_string())
204            })?;
205
206            // shape is typically [-1, seq_len, dim] or [-1, dim]
207            // Dynamic dimensions use -1, take the last positive dimension as embedding dimension
208            let dim = shape
209                .iter()
210                .rev()
211                .find_map(|&d| if d > 0 { Some(d as usize) } else { None })
212                .ok_or_else(|| {
213                    EmbeddingError::ParseError(format!(
214                        "Cannot infer embedding dimension from model output shape: {:?}",
215                        *shape
216                    ))
217                })?;
218
219            Ok(dim)
220        }
221
222        /// Simple whitespace tokenizer
223        ///
224        /// Splits text by whitespace into token ID sequence.
225        /// This is a basic implementation; production use should consider the `tokenizers` crate for subword tokenization.
226        fn simple_tokenize(text: &str) -> Vec<i64> {
227            // Simple whitespace tokenization, map each token's byte hash to an ID
228            text.split_whitespace()
229                .map(|word| {
230                    let mut h: u64 = 0xcbf29ce484222325;
231                    for b in word.bytes() {
232                        h ^= b as u64;
233                        h = h.wrapping_mul(0x100000001b3);
234                    }
235                    // Map to reasonable token ID range (0..30522 similar to BERT vocab size)
236                    (h % 30522) as i64
237                })
238                .collect()
239        }
240
241        /// Run ONNX inference, return raw output data
242        fn run_inference(
243            &self,
244            input_ids: &[i64],
245        ) -> Result<(Vec<usize>, Vec<f32>), EmbeddingError> {
246            let seq_len = input_ids.len();
247            if seq_len == 0 {
248                return Err(EmbeddingError::EmptyInput);
249            }
250
251            // Construct input_ids tensor: shape [1, seq_len]
252            let input_shape = vec![1i64, seq_len as i64];
253            let input_data = input_ids.to_vec();
254
255            let input_tensor = Tensor::from_array((input_shape, input_data)).map_err(|e| {
256                EmbeddingError::ApiError(format!("Failed to construct input tensor: {}", e))
257            })?;
258
259            // Get input name
260            let session = self.session.read().map_err(|e| {
261                EmbeddingError::ApiError(format!("Failed to acquire session read lock: {}", e))
262            })?;
263            let input_name = session
264                .inputs()
265                .first()
266                .map(|o| o.name().to_string())
267                .unwrap_or_else(|| "input_ids".to_string());
268
269            // run() requires &mut self, acquire write lock via RwLock
270            drop(session);
271            let mut session = self.session.write().map_err(|e| {
272                EmbeddingError::ApiError(format!("Failed to acquire session write lock: {}", e))
273            })?;
274            let outputs = session
275                .run(ort::inputs![input_name.as_str() => input_tensor]?)
276                .map_err(|e| EmbeddingError::ApiError(format!("ONNX inference failed: {}", e)))?;
277
278            // Get first output
279            let output_value = outputs.get(0).ok_or_else(|| {
280                EmbeddingError::ParseError("ONNX model has no output".to_string())
281            })?;
282
283            // Extract tensor data
284            let (shape, data) = output_value.try_extract_tensor::<f32>().map_err(|e| {
285                EmbeddingError::ParseError(format!("Failed to extract output tensor: {}", e))
286            })?;
287
288            let shape_vec: Vec<usize> = shape.iter().map(|&d| d as usize).collect();
289            let data_vec = data.to_vec();
290
291            Ok((shape_vec, data_vec))
292        }
293
294        /// Mean pooling: average over sequence dimension
295        ///
296        /// Input shape: [1, seq_len, dim] -> output: [dim]
297        /// or [1, dim] -> output: [dim]
298        fn mean_pool(shape: &[usize], data: &[f32]) -> Result<Vec<f32>, EmbeddingError> {
299            match shape.len() {
300                3 => {
301                    let dim = shape[2];
302                    let seq_len = shape[1];
303                    let mut result = vec![0.0f32; dim];
304
305                    for s in 0..seq_len {
306                        for d in 0..dim {
307                            result[d] += data[s * dim + d];
308                        }
309                    }
310
311                    for v in &mut result {
312                        *v /= seq_len as f32;
313                    }
314
315                    Ok(result)
316                }
317                2 => {
318                    let dim = shape[1];
319                    Ok(data[..dim].to_vec())
320                }
321                _ => Err(EmbeddingError::ParseError(format!(
322                    "Unsupported output dimension count: {}",
323                    shape.len()
324                ))),
325            }
326        }
327
328        /// L2 normalization
329        fn l2_normalize(vec: &mut [f32]) {
330            let norm: f32 = vec.iter().map(|x| x * x).sum::<f32>().sqrt();
331            if norm > 0.0 {
332                for v in vec.iter_mut() {
333                    *v /= norm;
334                }
335            }
336        }
337
338        /// Execute full embedding pipeline for a single text: tokenize -> inference -> mean pool -> L2 normalize
339        fn embed_single(&self, text: &str) -> Result<Vec<f32>, EmbeddingError> {
340            if text.trim().is_empty() {
341                return Err(EmbeddingError::EmptyInput);
342            }
343
344            let input_ids = Self::simple_tokenize(text);
345            if input_ids.is_empty() {
346                return Err(EmbeddingError::EmptyInput);
347            }
348
349            let (shape, raw_data) = self.run_inference(&input_ids)?;
350            let mut pooled = Self::mean_pool(&shape, &raw_data)?;
351            Self::l2_normalize(&mut pooled);
352            Ok(pooled)
353        }
354    }
355
356    #[async_trait]
357    impl Embeddings for LocalEmbeddings {
358        async fn embed_query(&self, text: &str) -> Result<Vec<f32>, EmbeddingError> {
359            // ONNX inference is CPU-intensive, run in blocking thread pool
360            let text = text.to_string();
361            tokio::task::spawn_blocking(move || self.embed_single(&text))
362                .await
363                .map_err(|e| EmbeddingError::ApiError(format!("Task execution failed: {}", e)))?
364        }
365
366        async fn embed_documents(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, EmbeddingError> {
367            if texts.is_empty() {
368                return Ok(Vec::new());
369            }
370
371            // Sequential inference (batch inference requires model support for multi-batch input)
372            let texts: Vec<String> = texts.iter().map(|s| s.to_string()).collect();
373            tokio::task::spawn_blocking(move || {
374                let mut results = Vec::with_capacity(texts.len());
375                for text in &texts {
376                    results.push(self.embed_single(text)?);
377                }
378                Ok(results)
379            })
380            .await
381            .map_err(|e| EmbeddingError::ApiError(format!("Task execution failed: {}", e)))?
382        }
383
384        fn dimension(&self) -> usize {
385            self.dim
386        }
387
388        fn model_name(&self) -> &str {
389            &self.model_name
390        }
391    }
392}
393
394// When local-embeddings feature is enabled, re-export LocalEmbeddings
395#[cfg(feature = "local-embeddings")]
396pub use nn::LocalEmbeddings;
397
398// ---------------------------------------------------------------------------
399// Backward compatibility: LocalEmbeddings without feature points to BagOfWordsEmbeddings
400// ---------------------------------------------------------------------------
401
402/// Without the `local-embeddings` feature, `LocalEmbeddings` is a type alias for `BagOfWordsEmbeddings`,
403/// maintaining backward compatibility.
404///
405/// With the `local-embeddings` feature enabled, `LocalEmbeddings` becomes the ONNX Runtime-based neural network implementation.
406#[cfg(not(feature = "local-embeddings"))]
407pub type LocalEmbeddings = BagOfWordsEmbeddings;
408
409// ---------------------------------------------------------------------------
410// Tests
411// ---------------------------------------------------------------------------
412
413#[cfg(test)]
414mod tests {
415    use super::*;
416    use crate::cosine_similarity;
417
418    // ---- BagOfWordsEmbeddings tests ----
419
420    #[tokio::test]
421    async fn test_bow_dimension() {
422        let e = BagOfWordsEmbeddings::new(128);
423        let v = e.embed_query("hello world").await.unwrap();
424        assert_eq!(v.len(), 128);
425        assert_eq!(e.dimension(), 128);
426    }
427
428    #[tokio::test]
429    async fn test_bow_same_text_same_vector() {
430        let e = BagOfWordsEmbeddings::new(64);
431        let a = e.embed_query("rust programming").await.unwrap();
432        let b = e.embed_query("rust programming").await.unwrap();
433        assert_eq!(a, b);
434    }
435
436    #[tokio::test]
437    async fn test_bow_different_text_different_vector() {
438        let e = BagOfWordsEmbeddings::new(64);
439        let a = e.embed_query("rust programming").await.unwrap();
440        let b = e.embed_query("cooking recipe pasta").await.unwrap();
441        assert_ne!(a, b);
442    }
443
444    #[tokio::test]
445    async fn test_bow_shared_words_more_similar() {
446        let e = BagOfWordsEmbeddings::new(256);
447        let base = e.embed_query("rust programming language").await.unwrap();
448        let similar = e.embed_query("rust programming tutorial").await.unwrap();
449        let different = e.embed_query("cooking pasta recipe").await.unwrap();
450
451        let sim_similar = cosine_similarity(&base, &similar).unwrap_or(0.0);
452        let sim_different = cosine_similarity(&base, &different).unwrap_or(0.0);
453        assert!(
454            sim_similar > sim_different,
455            "Shared words should be more similar: {} vs {}",
456            sim_similar,
457            sim_different
458        );
459    }
460
461    #[tokio::test]
462    async fn test_bow_normalized() {
463        let e = BagOfWordsEmbeddings::new(64);
464        let v = e.embed_query("some text here").await.unwrap();
465        let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
466        assert!((norm - 1.0).abs() < 1e-5, "norm = {}", norm);
467    }
468
469    #[tokio::test]
470    async fn test_bow_empty_text_returns_error() {
471        let e = BagOfWordsEmbeddings::new(64);
472        let result = e.embed_query("").await;
473        assert!(result.is_err());
474        assert!(matches!(result.unwrap_err(), EmbeddingError::EmptyInput));
475    }
476
477    #[tokio::test]
478    async fn test_bow_chinese_tokenize() {
479        let e = BagOfWordsEmbeddings::new(128);
480        let a = e.embed_query("机器学习").await.unwrap();
481        let b = e.embed_query("机器学习").await.unwrap();
482        assert_eq!(a, b);
483        let c = e.embed_query("深度学习").await.unwrap();
484        let sim = cosine_similarity(&a, &c).unwrap_or(0.0);
485        assert!(
486            sim > 0.0,
487            "Shared '学习' should have positive similarity: {}",
488            sim
489        );
490    }
491
492    #[test]
493    fn test_bow_tokenize_english() {
494        let t = BagOfWordsEmbeddings::tokenize("Hello, World! 123");
495        assert!(t.contains(&"hello".to_string()));
496        assert!(t.contains(&"world".to_string()));
497        assert!(t.contains(&"123".to_string()));
498    }
499
500    #[test]
501    fn test_bow_tokenize_chinese() {
502        let t = BagOfWordsEmbeddings::tokenize("机器学习");
503        assert!(t.contains(&"机".to_string()));
504        assert!(t.contains(&"学".to_string()));
505        assert_eq!(t.len(), 4);
506    }
507
508    #[test]
509    fn test_bow_model_name() {
510        let e = BagOfWordsEmbeddings::default_dim();
511        assert_eq!(e.model_name(), "local-bow");
512    }
513
514    // ---- LocalEmbeddings backward compatibility test (without feature, is BagOfWordsEmbeddings alias) ----
515
516    #[tokio::test]
517    async fn test_local_embeddings_backward_compat() {
518        // Without feature, LocalEmbeddings = BagOfWordsEmbeddings
519        let e = LocalEmbeddings::new(64);
520        let v = e.embed_query("test backward compat").await.unwrap();
521        assert_eq!(v.len(), 64);
522        assert_eq!(e.model_name(), "local-bow");
523    }
524
525    // ---- ONNX LocalEmbeddings tests (requires local-embeddings feature) ----
526
527    #[cfg(feature = "local-embeddings")]
528    mod nn_tests {
529        use super::*;
530
531        #[test]
532        fn test_l2_normalize() {
533            let mut v = vec![3.0, 4.0];
534            LocalEmbeddings::l2_normalize(&mut v);
535            let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
536            assert!((norm - 1.0).abs() < 1e-5);
537            assert!((v[0] - 0.6).abs() < 1e-5);
538            assert!((v[1] - 0.8).abs() < 1e-5);
539        }
540
541        #[test]
542        fn test_l2_normalize_zero() {
543            let mut v = vec![0.0, 0.0, 0.0];
544            LocalEmbeddings::l2_normalize(&mut v);
545            assert!(v.iter().all(|x| *x == 0.0));
546        }
547
548        #[test]
549        fn test_mean_pool_3d() {
550            // shape [1, 2, 3]: 2 tokens, 3 dimensions
551            let shape = vec![1usize, 2, 3];
552            let data = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
553            let result = LocalEmbeddings::mean_pool(&shape, &data).unwrap();
554            assert_eq!(result.len(), 3);
555            assert!((result[0] - 2.5).abs() < 1e-5);
556            assert!((result[1] - 3.5).abs() < 1e-5);
557            assert!((result[2] - 4.5).abs() < 1e-5);
558        }
559
560        #[test]
561        fn test_mean_pool_2d() {
562            // shape [1, 3]: direct extraction
563            let shape = vec![1usize, 3];
564            let data = vec![1.0, 2.0, 3.0];
565            let result = LocalEmbeddings::mean_pool(&shape, &data).unwrap();
566            assert_eq!(result.len(), 3);
567            assert!((result[0] - 1.0).abs() < 1e-5);
568            assert!((result[1] - 2.0).abs() < 1e-5);
569            assert!((result[2] - 3.0).abs() < 1e-5);
570        }
571
572        #[test]
573        fn test_simple_tokenize() {
574            let tokens = LocalEmbeddings::simple_tokenize("hello world test");
575            assert_eq!(tokens.len(), 3);
576            // Same word should produce same token ID
577            let tokens2 = LocalEmbeddings::simple_tokenize("hello");
578            assert_eq!(tokens[0], tokens2[0]);
579        }
580
581        #[test]
582        fn test_simple_tokenize_empty() {
583            let tokens = LocalEmbeddings::simple_tokenize("");
584            assert!(tokens.is_empty());
585        }
586    }
587}