navi_core/memory/
embedding.rs1use anyhow::Result;
11use std::path::PathBuf;
12use std::sync::{Arc, OnceLock};
13
14pub const EMBED_DIM: usize = 256;
17
18pub const FULL_EMBED_DIM: usize = 1024;
20
21pub const DEFAULT_MODEL_REPO: &str = "Qwen/Qwen3-Embedding-0.6B-GGUF";
23pub const DEFAULT_MODEL_FILE: &str = "Qwen3-Embedding-0.6B-Q8_0.gguf";
24pub const DEFAULT_TOKENIZER_REPO: &str = "Qwen/Qwen3-Embedding-0.6B";
25pub const DEFAULT_TOKENIZER_FILE: &str = "tokenizer.json";
26
27static EMBEDDER_CACHE: OnceLock<std::sync::Mutex<Option<Arc<dyn Embedder>>>> = OnceLock::new();
29
30pub fn get_cached_embedder(
33 model_path: &PathBuf,
34 tokenizer_path: &PathBuf,
35) -> Option<Arc<dyn Embedder>> {
36 let cache = EMBEDDER_CACHE.get_or_init(|| std::sync::Mutex::new(None));
37 let mut guard = cache.lock().ok()?;
38
39 if let Some(ref embedder) = *guard {
40 return Some(embedder.clone());
41 }
42
43 let config = EmbeddingConfig {
45 model_path: model_path.clone(),
46 tokenizer_path: tokenizer_path.clone(),
47 ..Default::default()
48 };
49 let embedder = create_embedder(config);
50 let embedder_arc: Arc<dyn Embedder> = Arc::from(embedder);
51
52 match embedder_arc.embed("test") {
55 Ok(_) => {
56 *guard = Some(embedder_arc.clone());
57 Some(embedder_arc)
58 }
59 Err(_) => None,
60 }
61}
62
63pub trait Embedder: Send + Sync {
65 fn embed(&self, text: &str) -> Result<Vec<f32>>;
68}
69
70pub struct NoEmbedder;
73
74impl Embedder for NoEmbedder {
75 fn embed(&self, _text: &str) -> Result<Vec<f32>> {
76 anyhow::bail!("embeddings feature is not enabled")
77 }
78}
79
80#[derive(Debug, Clone)]
82pub struct EmbeddingConfig {
83 pub model_path: PathBuf,
85 pub tokenizer_path: PathBuf,
87 pub normalize: bool,
89}
90
91impl Default for EmbeddingConfig {
92 fn default() -> Self {
93 Self {
94 model_path: PathBuf::new(),
95 tokenizer_path: PathBuf::new(),
96 normalize: true,
97 }
98 }
99}
100
101#[cfg(feature = "embeddings")]
102mod candle_embedder {
103 use super::*;
104 use anyhow::{Context, Result as AnyResult};
105 use candle_core::quantized::gguf_file;
106 use candle_core::{DType, Device, Tensor};
107 use candle_transformers::models::quantized_qwen2::ModelWeights;
108 use std::fs::File;
109 use std::io::BufReader;
110 use tokenizers::Tokenizer;
111
112 pub struct CandleEmbedder {
115 model: std::sync::Mutex<ModelWeights>,
116 tokenizer: Tokenizer,
117 device: Device,
118 config: EmbeddingConfig,
119 }
120
121 impl CandleEmbedder {
122 pub fn load(config: EmbeddingConfig) -> AnyResult<Self> {
124 let device = Device::Cpu;
125
126 let file = File::open(&config.model_path)
128 .with_context(|| format!("Failed to open GGUF file: {:?}", config.model_path))?;
129 let mut reader = BufReader::new(file);
130
131 let ct =
133 gguf_file::Content::read(&mut reader).context("Failed to read GGUF content")?;
134
135 let model = ModelWeights::from_gguf(ct, &mut reader, &device)
137 .context("Failed to build quantized Qwen2 model from GGUF")?;
138
139 let tokenizer = Tokenizer::from_file(&config.tokenizer_path).map_err(|e| {
141 anyhow::anyhow!(
142 "Failed to load tokenizer from {:?}: {}",
143 config.tokenizer_path,
144 e
145 )
146 })?;
147
148 Ok(Self {
149 model: std::sync::Mutex::new(model),
150 tokenizer,
151 device,
152 config,
153 })
154 }
155
156 fn mean_pool(
158 &self,
159 token_embeddings: &Tensor,
160 attention_mask: &Tensor,
161 ) -> AnyResult<Tensor> {
162 let mask = attention_mask.to_dtype(DType::F32)?.unsqueeze(2)?; let masked = token_embeddings.broadcast_mul(&mask)?;
165 let sum = masked.sum(1)?; let mask_sum = mask.sum(1)?; let pooled = sum.broadcast_div(&mask_sum)?;
169
170 if self.config.normalize {
171 let norm = pooled.sqr()?.sum(1)?.sqrt()?;
172 let pooled = pooled.broadcast_div(&norm.unsqueeze(1)?)?;
173 Ok(pooled)
174 } else {
175 Ok(pooled)
176 }
177 }
178 }
179
180 impl Embedder for CandleEmbedder {
181 fn embed(&self, text: &str) -> AnyResult<Vec<f32>> {
182 let encoding = self
183 .tokenizer
184 .encode(text, true)
185 .map_err(|e| anyhow::anyhow!("Tokenization failed: {}", e))?;
186
187 let input_ids = encoding.get_ids();
188 let attention_mask = encoding.get_attention_mask();
189
190 let input_ids_tensor = Tensor::from_slice(
192 input_ids
193 .iter()
194 .map(|&v| v as u32)
195 .collect::<Vec<_>>()
196 .as_slice(),
197 (1, input_ids.len()),
198 &self.device,
199 )?;
200
201 let attention_mask_tensor = Tensor::from_slice(
203 attention_mask
204 .iter()
205 .map(|&v| v as u32)
206 .collect::<Vec<_>>()
207 .as_slice(),
208 (1, attention_mask.len()),
209 &self.device,
210 )?;
211
212 let mut model = self
214 .model
215 .lock()
216 .map_err(|e| anyhow::anyhow!("model lock poisoned: {}", e))?;
217 let embedded = model.forward(&input_ids_tensor, 0)?;
218 let pooled = self.mean_pool(&embedded, &attention_mask_tensor)?;
222 let full_embedding = pooled
226 .to_vec2::<f32>()?
227 .into_iter()
228 .next()
229 .unwrap_or_default();
230
231 let truncated: Vec<f32> = full_embedding.into_iter().take(EMBED_DIM).collect();
233
234 if self.config.normalize && !truncated.is_empty() {
236 let norm: f32 = truncated.iter().map(|v| v * v).sum::<f32>().sqrt();
237 if norm > 0.0 {
238 return Ok(truncated.iter().map(|v| v / norm).collect());
239 }
240 }
241
242 Ok(truncated)
243 }
244 }
245}
246
247#[cfg(feature = "embeddings")]
248pub use candle_embedder::CandleEmbedder;
249
250#[cfg(feature = "embeddings")]
252pub fn create_embedder(config: EmbeddingConfig) -> Box<dyn Embedder> {
253 if config.model_path.exists() && config.tokenizer_path.exists() {
254 match CandleEmbedder::load(config) {
255 Ok(embedder) => {
256 tracing::info!("Local embedding model loaded successfully");
257 return Box::new(embedder);
258 }
259 Err(e) => {
260 tracing::warn!(
261 "Failed to load embedding model: {}, falling back to text search",
262 e
263 );
264 }
265 }
266 } else {
267 if !config.model_path.exists() {
268 tracing::info!(
269 "Embedding model not found at {:?}. Download with: huggingface-cli download {} {}",
270 config.model_path,
271 DEFAULT_MODEL_REPO,
272 DEFAULT_MODEL_FILE
273 );
274 }
275 if !config.tokenizer_path.exists() {
276 tracing::info!(
277 "Tokenizer not found at {:?}. Download with: huggingface-cli download {} {}",
278 config.tokenizer_path,
279 DEFAULT_TOKENIZER_REPO,
280 DEFAULT_TOKENIZER_FILE
281 );
282 }
283 }
284
285 Box::new(NoEmbedder)
286}
287
288#[cfg(not(feature = "embeddings"))]
290pub fn create_embedder(_config: EmbeddingConfig) -> Box<dyn Embedder> {
291 Box::new(NoEmbedder)
292}
293
294#[cfg(feature = "embeddings")]
296pub fn embeddings_available() -> bool {
297 true
298}
299
300#[cfg(not(feature = "embeddings"))]
302pub fn embeddings_available() -> bool {
303 false
304}
305
306#[cfg(test)]
307mod tests {
308 use super::*;
309
310 #[test]
311 fn test_no_embedder_returns_error() {
312 let embedder = NoEmbedder;
313 assert!(embedder.embed("test").is_err());
314 }
315
316 #[test]
317 fn test_embed_dim() {
318 assert_eq!(EMBED_DIM, 256);
319 assert_eq!(FULL_EMBED_DIM, 1024);
320 }
321
322 #[test]
323 fn test_create_embedder_without_model_file() {
324 let config = EmbeddingConfig::default();
325 let embedder = create_embedder(config);
326 assert!(embedder.embed("test").is_err());
328 }
329}