Skip to main content

gonzalo_embed/
lib.rs

1//! Local CPU sentence-embedding [`Embedder`](gonzalo_vector::Embedder) for
2//! gonzalo, built on Candle + `all-MiniLM-L6-v2` (ADR 0013).
3//!
4//! [`CandleEmbedder::load`] resolves the model weights (via an
5//! [`EmbedderConfig::model_path`] override, else a one-time anonymous `hf-hub`
6//! download) and loads the BERT model + tokenizer once on CPU. [`embed`] then
7//! tokenizes, runs the forward pass, masked-mean-pools the token states, and
8//! L2-normalizes to a 384-dim unit vector. The synchronous CPU forward runs
9//! inside `spawn_blocking` so it never blocks the async runtime.
10//!
11//! [`embed`]: CandleEmbedder::embed
12
13use 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/// Configuration for [`CandleEmbedder::load`].
25#[derive(Debug, Clone)]
26pub struct EmbedderConfig {
27    /// HuggingFace model id. Default `sentence-transformers/all-MiniLM-L6-v2`.
28    pub model_id: String,
29    /// Model revision (branch, tag, or commit). Default `main`.
30    pub revision: String,
31    /// If set, load `model.safetensors`/`tokenizer.json`/`config.json` from this
32    /// directory instead of downloading — fully offline.
33    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/// A local CPU sentence embedder (Candle + all-MiniLM). Cheap to clone.
53#[derive(Clone)]
54pub struct CandleEmbedder {
55    inner: Arc<Inner>,
56}
57
58impl CandleEmbedder {
59    /// Resolve the weights, tokenizer, and config, then load the model on CPU.
60    /// All failures surface as [`CoreError::Backend`].
61    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        // Safe (non-mmap) load — the workspace forbids `unsafe`.
90        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    /// The synchronous embedding pipeline: tokenize → forward → masked mean-pool
106    /// → L2-normalize → a 384-dim unit vector.
107    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
142/// Masked mean-pooling: average the per-token hidden states `(batch, seq,
143/// hidden)` over the sequence, counting only positions where `attention_mask`
144/// `(batch, seq)` is 1. Padding/masked positions are excluded.
145fn mean_pool(hidden: &Tensor, attention_mask: &Tensor) -> candle_core::Result<Tensor> {
146    let mask = attention_mask.to_dtype(DType::F32)?.unsqueeze(2)?; // (b, s, 1)
147    let summed = hidden.broadcast_mul(&mask)?.sum(1)?; // (b, h)
148    let counts = mask.sum(1)?; // (b, 1)
149    summed.broadcast_div(&counts)
150}
151
152/// L2-normalize each row of `(batch, hidden)` to unit length.
153///
154/// The norm is floored at a tiny epsilon so a zero/degenerate pooled row
155/// divides by a small positive value instead of `0.0`, yielding a finite
156/// (near-zero) vector rather than emitting `NaN`.
157fn l2_normalize(v: &Tensor) -> candle_core::Result<Tensor> {
158    let norm = v.sqr()?.sum_keepdim(D::Minus1)?.sqrt()?.maximum(1e-12f64)?; // (b, 1)
159    v.broadcast_div(&norm)
160}
161
162/// Map any backend/model/tokenizer/IO error into [`CoreError::Backend`].
163fn 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        // A 3-4-5 right triangle row: (3, 4) normalizes to (0.6, 0.8).
174        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        // A degenerate all-zero pooled row must not divide by zero and emit NaN.
190        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        // Two real tokens [1,1] and [3,3] plus a padding token [100,100] that
203        // the mask (1,1,0) must exclude. Mean of the real tokens is [2,2].
204        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    // Real-model check — excluded from default `cargo test` (downloads ~90MB on
222    // first run). Run with `cargo test -p gonzalo-embed -- --ignored`.
223    #[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}