//! Optional MiniLM sentence embedder for the LongMemEval bench. //! //! Compiled only under the `embeddings` feature, so the default build of a //! project that prides itself on having no heavyweight dependencies stays //! exactly as it was. Without it the bench runs BM25-only, as it always has. //! //! Loads `sentence-transformers/all-MiniLM-L6-v2` — the same checkpoint //! omni-cortex uses — and produces 384-d mean-pooled, L2-normalised sentence //! embeddings, which is the published recipe for this model (mean over token //! states weighted by the attention mask, *not* the `[CLS]` pooler output). use std::collections::HashMap; use std::path::Path; use candle_core::{DType, Device, Tensor}; use candle_nn::VarBuilder; use candle_transformers::models::bert::{BertModel, Config, HiddenAct}; use tokenizers::Tokenizer; /// Sequences encoded per forward pass. Larger batches amortise the transformer /// call; 64 keeps peak memory modest while still saturating a CPU. const BATCH: usize = 64; /// A loaded MiniLM encoder. pub struct Embedder { model: BertModel, tokenizer: Tokenizer, device: Device, } impl Embedder { /// Load from a directory holding `model.safetensors` and `tokenizer.json`. /// /// `config.json` is read when present; otherwise the published MiniLM-L6-v2 /// architecture constants are used, which are pinned rather than guessed. pub fn load(dir: &Path) -> Result> { // CUDA when the feature is on and a device is actually present; the CPU // path is correct but roughly two orders of magnitude slower, which is // the difference between minutes and most of a day on the full haystack. let device = match Device::new_cuda(0) { Ok(d) => { eprintln!("Embedder: CUDA device 0"); d } Err(e) => { // Loud, because the CPU path is correct but ~100x slower: the // full longmemeval_s haystack is minutes on a GPU and most of a // day on 8 cores. Silently falling back looks like a hang. eprintln!("Embedder: CPU — CUDA unavailable ({e})"); eprintln!( " WARNING: CPU embedding is roughly two orders of magnitude slower.\n Expect minutes for longmemeval_oracle and many hours for the full\n longmemeval_s haystack. For the GPU path, rebuild with\n `--features embeddings-cuda` and make sure `nvcc` is on PATH\n (it ships in /usr/local/cuda/bin, which is often not exported)." ); Device::Cpu } }; let weights = dir.join("model.safetensors"); let tok_path = dir.join("tokenizer.json"); let config: Config = match std::fs::read_to_string(dir.join("config.json")) { Ok(raw) => serde_json::from_str(&raw)?, Err(_) => Config { vocab_size: 30_522, hidden_size: 384, num_hidden_layers: 6, num_attention_heads: 12, intermediate_size: 1_536, hidden_act: HiddenAct::Gelu, hidden_dropout_prob: 0.0, max_position_embeddings: 512, type_vocab_size: 2, initializer_range: 0.02, layer_norm_eps: 1e-12, pad_token_id: 0, position_embedding_type: Default::default(), use_cache: false, classifier_dropout: None, model_type: None, }, }; let vb = unsafe { VarBuilder::from_mmaped_safetensors(&[weights], DType::F32, &device)? }; let model = BertModel::load(vb, &config)?; let tokenizer = Tokenizer::from_file(&tok_path).map_err(|e| e.to_string())?; Ok(Self { model, tokenizer, device, }) } /// Encode `texts` into 384-d unit vectors, in order. fn encode_batch(&self, texts: &[&str]) -> Result>, Box> { let mut tk = self.tokenizer.clone(); let tk = tk .with_padding(Some(tokenizers::PaddingParams::default())) .with_truncation(Some(tokenizers::TruncationParams { max_length: 512, ..Default::default() })) .map_err(|e| e.to_string())?; let encodings = tk .encode_batch(texts.to_vec(), true) .map_err(|e| e.to_string())?; let ids: Vec = encodings .iter() .flat_map(|e| e.get_ids().to_vec()) .collect(); let mask: Vec = encodings .iter() .flat_map(|e| e.get_attention_mask().to_vec()) .collect(); let (b, l) = (encodings.len(), encodings[0].get_ids().len()); let ids = Tensor::from_vec(ids, (b, l), &self.device)?; let mask = Tensor::from_vec(mask, (b, l), &self.device)?; let type_ids = ids.zeros_like()?; let hidden = self.model.forward(&ids, &type_ids, Some(&mask))?; // Mean-pool over real tokens only: sum(hidden * mask) / sum(mask). let mask_f = mask.to_dtype(DType::F32)?.unsqueeze(2)?; let summed = hidden.broadcast_mul(&mask_f)?.sum(1)?; let counts = mask_f.sum(1)?.clamp(1e-9, f32::INFINITY)?; let pooled = summed.broadcast_div(&counts)?; // L2-normalise so cosine similarity is a plain dot product. let norm = pooled .sqr()? .sum_keepdim(1)? .sqrt()? .clamp(1e-12, f32::INFINITY)?; let normed = pooled.broadcast_div(&norm)?; Ok(normed.to_vec2::()?) } /// Encode every distinct string in `texts` once, returning a lookup map. /// /// LongMemEval's haystack sessions are drawn from a shared pool, so the same /// turn text recurs across many questions. Deduplicating before encoding is /// the difference between encoding the corpus once and encoding it per /// question. pub fn encode_unique( &self, texts: impl IntoIterator, ) -> Result>, Box> { let mut unique: Vec = texts.into_iter().collect(); unique.sort_unstable(); unique.dedup(); let total = unique.len(); eprintln!("Embedding {total} unique texts with MiniLM (batch {BATCH})..."); let mut out = HashMap::with_capacity(total); for (n, chunk) in unique.chunks(BATCH).enumerate() { let refs: Vec<&str> = chunk.iter().map(String::as_str).collect(); let vecs = self.encode_batch(&refs)?; for (text, v) in chunk.iter().zip(vecs) { out.insert(text.clone(), v); } if n % 50 == 0 { eprint!("\r [{}/{}] embedded...", (n * BATCH).min(total), total); } } eprintln!("\r [{total}/{total}] embedded. "); Ok(out) } }