1use std::path::PathBuf;
14use std::sync::Arc;
15
16use async_trait::async_trait;
17use candle_core::{D, DType, Device, Tensor};
18use candle_nn::VarBuilder;
19use candle_transformers::models::bert::{BertModel, Config, DTYPE};
20use gonzalo_core::{CoreError, Result};
21use gonzalo_vector::Embedder;
22use tokenizers::Tokenizer;
23
24#[derive(Debug, Clone)]
26pub struct EmbedderConfig {
27 pub model_id: String,
29 pub revision: String,
31 pub model_path: Option<PathBuf>,
34}
35
36impl Default for EmbedderConfig {
37 fn default() -> Self {
38 Self {
39 model_id: "sentence-transformers/all-MiniLM-L6-v2".to_string(),
40 revision: "main".to_string(),
41 model_path: None,
42 }
43 }
44}
45
46struct Inner {
47 tokenizer: Tokenizer,
48 model: BertModel,
49 device: Device,
50}
51
52#[derive(Clone)]
54pub struct CandleEmbedder {
55 inner: Arc<Inner>,
56}
57
58impl CandleEmbedder {
59 pub async fn load(config: EmbedderConfig) -> Result<Self> {
62 let (weights, tokenizer_path, config_path) = match &config.model_path {
63 Some(dir) => (
64 dir.join("model.safetensors"),
65 dir.join("tokenizer.json"),
66 dir.join("config.json"),
67 ),
68 None => {
69 let api = hf_hub::api::tokio::ApiBuilder::new()
70 .build()
71 .map_err(backend)?;
72 let repo = api.repo(hf_hub::Repo::with_revision(
73 config.model_id.clone(),
74 hf_hub::RepoType::Model,
75 config.revision.clone(),
76 ));
77 (
78 repo.get("model.safetensors").await.map_err(backend)?,
79 repo.get("tokenizer.json").await.map_err(backend)?,
80 repo.get("config.json").await.map_err(backend)?,
81 )
82 }
83 };
84
85 let device = Device::Cpu;
86 let tokenizer = Tokenizer::from_file(&tokenizer_path).map_err(backend)?;
87 let cfg: Config = serde_json::from_slice(&std::fs::read(&config_path).map_err(backend)?)
88 .map_err(backend)?;
89 let tensors = candle_core::safetensors::load(&weights, &device).map_err(backend)?;
91 let vb = VarBuilder::from_tensors(tensors, DTYPE, &device);
92 let model = BertModel::load(vb, &cfg).map_err(backend)?;
93
94 Ok(Self {
95 inner: Arc::new(Inner {
96 tokenizer,
97 model,
98 device,
99 }),
100 })
101 }
102}
103
104impl Inner {
105 fn embed_blocking(&self, text: &str) -> Result<Vec<f32>> {
108 let encoding = self.tokenizer.encode(text, true).map_err(backend)?;
109 let ids = Tensor::new(encoding.get_ids(), &self.device)
110 .and_then(|t| t.unsqueeze(0))
111 .map_err(backend)?;
112 let attention_mask = Tensor::new(encoding.get_attention_mask(), &self.device)
113 .and_then(|t| t.unsqueeze(0))
114 .map_err(backend)?;
115 let token_type_ids = ids.zeros_like().map_err(backend)?;
116
117 let hidden = self
118 .model
119 .forward(&ids, &token_type_ids, Some(&attention_mask))
120 .map_err(backend)?;
121
122 let pooled = mean_pool(&hidden, &attention_mask).map_err(backend)?;
123 let normalized = l2_normalize(&pooled).map_err(backend)?;
124 normalized
125 .squeeze(0)
126 .and_then(|t| t.to_vec1::<f32>())
127 .map_err(backend)
128 }
129}
130
131#[async_trait]
132impl Embedder for CandleEmbedder {
133 async fn embed(&self, text: &str) -> Result<Vec<f32>> {
134 let inner = self.inner.clone();
135 let text = text.to_string();
136 tokio::task::spawn_blocking(move || inner.embed_blocking(&text))
137 .await
138 .map_err(backend)?
139 }
140}
141
142fn mean_pool(hidden: &Tensor, attention_mask: &Tensor) -> candle_core::Result<Tensor> {
146 let mask = attention_mask.to_dtype(DType::F32)?.unsqueeze(2)?; let summed = hidden.broadcast_mul(&mask)?.sum(1)?; let counts = mask.sum(1)?; summed.broadcast_div(&counts)
150}
151
152fn l2_normalize(v: &Tensor) -> candle_core::Result<Tensor> {
158 let norm = v.sqr()?.sum_keepdim(D::Minus1)?.sqrt()?.maximum(1e-12f64)?; v.broadcast_div(&norm)
160}
161
162fn backend<E: std::fmt::Display>(e: E) -> CoreError {
164 CoreError::Backend(e.to_string())
165}
166
167#[cfg(test)]
168mod tests {
169 use super::*;
170
171 #[test]
172 fn l2_normalize_yields_unit_length() {
173 let v = Tensor::from_vec(vec![3f32, 4.0], (1, 2), &Device::Cpu).unwrap();
175 let out = l2_normalize(&v)
176 .unwrap()
177 .squeeze(0)
178 .unwrap()
179 .to_vec1::<f32>()
180 .unwrap();
181 assert!((out[0] - 0.6).abs() < 1e-5);
182 assert!((out[1] - 0.8).abs() < 1e-5);
183 let len = (out[0] * out[0] + out[1] * out[1]).sqrt();
184 assert!((len - 1.0).abs() < 1e-5);
185 }
186
187 #[test]
188 fn l2_normalize_zero_row_is_finite_not_nan() {
189 let v = Tensor::from_vec(vec![0f32, 0.0], (1, 2), &Device::Cpu).unwrap();
191 let out = l2_normalize(&v)
192 .unwrap()
193 .squeeze(0)
194 .unwrap()
195 .to_vec1::<f32>()
196 .unwrap();
197 assert!(out.iter().all(|x| x.is_finite()), "output must be finite");
198 }
199
200 #[test]
201 fn mean_pool_ignores_masked_positions() {
202 let hidden = Tensor::from_vec(
205 vec![1f32, 1.0, 3.0, 3.0, 100.0, 100.0],
206 (1, 3, 2),
207 &Device::Cpu,
208 )
209 .unwrap();
210 let mask = Tensor::from_vec(vec![1u32, 1, 0], (1, 3), &Device::Cpu).unwrap();
211 let out = mean_pool(&hidden, &mask)
212 .unwrap()
213 .squeeze(0)
214 .unwrap()
215 .to_vec1::<f32>()
216 .unwrap();
217 assert!((out[0] - 2.0).abs() < 1e-5);
218 assert!((out[1] - 2.0).abs() < 1e-5);
219 }
220
221 #[tokio::test]
224 #[ignore = "downloads the all-MiniLM model on first run"]
225 async fn real_model_embeds_and_ranks_semantically() {
226 let embedder = CandleEmbedder::load(EmbedderConfig::default())
227 .await
228 .unwrap();
229 let anchor = embedder.embed("the cat sat on the mat").await.unwrap();
230 assert_eq!(anchor.len(), 384, "all-MiniLM produces 384-dim vectors");
231 let len: f32 = anchor.iter().map(|x| x * x).sum::<f32>().sqrt();
232 assert!((len - 1.0).abs() < 1e-3, "output is L2-normalized");
233
234 let related = embedder.embed("a kitten naps on a rug").await.unwrap();
235 let unrelated = embedder.embed("the diesel engine roared").await.unwrap();
236 let cos = |a: &[f32], b: &[f32]| a.iter().zip(b).map(|(x, y)| x * y).sum::<f32>();
237 assert!(
238 cos(&anchor, &related) > cos(&anchor, &unrelated),
239 "a semantically related sentence should score higher"
240 );
241 }
242}