390 lines
13 KiB
Rust
390 lines
13 KiB
Rust
use std::sync::Mutex;
|
|
use std::sync::OnceLock;
|
|
|
|
use candle_core::{DType, Device, Tensor};
|
|
use candle_nn::VarBuilder;
|
|
use candle_transformers::models::bert::{BertModel, Config};
|
|
use tokenizers::Tokenizer;
|
|
|
|
struct CandleEmbeddingModel {
|
|
model: BertModel,
|
|
tokenizer: Tokenizer,
|
|
device: Device,
|
|
}
|
|
|
|
impl CandleEmbeddingModel {
|
|
fn new() -> Result<Self, String> {
|
|
if cfg!(test) || std::env::var("MCP_OFFLINE_EMBEDDINGS").is_ok() {
|
|
return Err("Offline embedding mode active (tests/offline flag)".to_string());
|
|
}
|
|
let client = hf_hub::HFClientSync::new().map_err(|e| e.to_string())?;
|
|
let repo = client.model("sentence-transformers", "all-MiniLM-L6-v2");
|
|
|
|
let config_file = repo
|
|
.download_file()
|
|
.filename("config.json")
|
|
.send()
|
|
.map_err(|e| format!("Failed to download config.json: {}", e))?;
|
|
let tokenizer_file = repo
|
|
.download_file()
|
|
.filename("tokenizer.json")
|
|
.send()
|
|
.map_err(|e| format!("Failed to download tokenizer.json: {}", e))?;
|
|
let weights_file = repo
|
|
.download_file()
|
|
.filename("model.safetensors")
|
|
.send()
|
|
.map_err(|e| format!("Failed to download model.safetensors: {}", e))?;
|
|
|
|
let config_str = std::fs::read_to_string(&config_file)
|
|
.map_err(|e| format!("Failed to read config.json: {}", e))?;
|
|
let config: Config = serde_json::from_str(&config_str)
|
|
.map_err(|e| format!("Failed to parse config.json: {}", e))?;
|
|
|
|
let mut tokenizer = Tokenizer::from_file(&tokenizer_file)
|
|
.map_err(|e| format!("Failed to load tokenizer: {}", e))?;
|
|
tokenizer.with_padding(Some(tokenizers::PaddingParams::default()));
|
|
|
|
let device = Device::Cpu;
|
|
let vb = unsafe {
|
|
VarBuilder::from_mmaped_safetensors(&[weights_file], DType::F32, &device)
|
|
.map_err(|e| format!("Failed to load safetensors: {}", e))?
|
|
};
|
|
|
|
let model =
|
|
BertModel::load(vb, &config).map_err(|e| format!("Failed to load BertModel: {}", e))?;
|
|
|
|
Ok(Self {
|
|
model,
|
|
tokenizer,
|
|
device,
|
|
})
|
|
}
|
|
|
|
fn embed(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, String> {
|
|
if texts.is_empty() {
|
|
return Ok(Vec::new());
|
|
}
|
|
|
|
let encodings = self
|
|
.tokenizer
|
|
.encode_batch(texts.to_vec(), true)
|
|
.map_err(|e| format!("Failed to encode texts: {}", e))?;
|
|
|
|
let batch_size = encodings.len();
|
|
if batch_size == 0 {
|
|
return Ok(Vec::new());
|
|
}
|
|
let seq_len = encodings[0].get_ids().len();
|
|
if seq_len == 0 {
|
|
return Ok(vec![vec![0.0; 384]; batch_size]);
|
|
}
|
|
|
|
let mut all_ids: Vec<u32> = Vec::with_capacity(batch_size * seq_len);
|
|
let mut all_type_ids: Vec<u32> = Vec::with_capacity(batch_size * seq_len);
|
|
let mut all_attention_mask: Vec<u32> = Vec::with_capacity(batch_size * seq_len);
|
|
|
|
for enc in &encodings {
|
|
all_ids.extend(enc.get_ids());
|
|
all_type_ids.extend(enc.get_type_ids());
|
|
all_attention_mask.extend(enc.get_attention_mask());
|
|
}
|
|
|
|
let input_ids = Tensor::from_vec(all_ids, (batch_size, seq_len), &self.device)
|
|
.map_err(|e| format!("Failed to build input_ids tensor: {}", e))?;
|
|
let token_type_ids = Tensor::from_vec(all_type_ids, (batch_size, seq_len), &self.device)
|
|
.map_err(|e| format!("Failed to build token_type_ids tensor: {}", e))?;
|
|
let attention_mask =
|
|
Tensor::from_vec(all_attention_mask, (batch_size, seq_len), &self.device)
|
|
.map_err(|e| format!("Failed to build attention_mask tensor: {}", e))?;
|
|
|
|
let sequence_output = self
|
|
.model
|
|
.forward(&input_ids, &token_type_ids, Some(&attention_mask))
|
|
.map_err(|e| format!("Bert forward failed: {}", e))?;
|
|
|
|
// Mean pooling: sum(sequence_output * mask) / clamp(sum(mask), min=1e-9)
|
|
let mask_f32 = attention_mask
|
|
.to_dtype(DType::F32)
|
|
.map_err(|e| e.to_string())?
|
|
.unsqueeze(2)
|
|
.map_err(|e| e.to_string())?;
|
|
|
|
let sum_embeddings = sequence_output
|
|
.broadcast_mul(&mask_f32)
|
|
.map_err(|e| e.to_string())?
|
|
.sum(1)
|
|
.map_err(|e| e.to_string())?;
|
|
|
|
let sum_mask = mask_f32
|
|
.sum(1)
|
|
.map_err(|e| e.to_string())?
|
|
.clamp(1e-9, f32::MAX)
|
|
.map_err(|e| e.to_string())?;
|
|
|
|
let mean_pooled = sum_embeddings
|
|
.broadcast_div(&sum_mask)
|
|
.map_err(|e| e.to_string())?;
|
|
|
|
// L2 Normalization
|
|
let norm = mean_pooled
|
|
.sqr()
|
|
.map_err(|e| e.to_string())?
|
|
.sum_keepdim(1)
|
|
.map_err(|e| e.to_string())?
|
|
.sqrt()
|
|
.map_err(|e| e.to_string())?;
|
|
|
|
let normalized = mean_pooled
|
|
.broadcast_div(&norm)
|
|
.map_err(|e| e.to_string())?;
|
|
|
|
normalized.to_vec2::<f32>().map_err(|e| e.to_string())
|
|
}
|
|
}
|
|
|
|
fn fallback_embed(text: &str) -> Vec<f32> {
|
|
const DIM: usize = 384;
|
|
let mut vec = vec![0.0f32; DIM];
|
|
let words: Vec<&str> = text.split_whitespace().collect();
|
|
if words.is_empty() {
|
|
vec[0] = 1.0;
|
|
return vec;
|
|
}
|
|
|
|
for word in words {
|
|
let clean: String = word
|
|
.chars()
|
|
.filter(|c| c.is_alphanumeric())
|
|
.flat_map(|c| c.to_lowercase())
|
|
.collect();
|
|
if clean.is_empty() {
|
|
continue;
|
|
}
|
|
|
|
let hash_bytes = blake3::hash(clean.as_bytes());
|
|
let h = u64::from_le_bytes(hash_bytes.as_bytes()[0..8].try_into().unwrap());
|
|
let idx = (h as usize) % DIM;
|
|
let sign = if (h >> 32) & 1 == 0 { 1.0f32 } else { -1.0f32 };
|
|
vec[idx] += sign;
|
|
|
|
let chars: Vec<char> = clean.chars().collect();
|
|
for window in chars.windows(3) {
|
|
let window_str: String = window.iter().collect();
|
|
let h2_bytes = blake3::hash(window_str.as_bytes());
|
|
let hv = u64::from_le_bytes(h2_bytes.as_bytes()[0..8].try_into().unwrap());
|
|
let idx2 = (hv as usize) % DIM;
|
|
let s2 = if (hv >> 32) & 1 == 0 { 0.5f32 } else { -0.5f32 };
|
|
vec[idx2] += s2;
|
|
}
|
|
}
|
|
|
|
let norm_sq: f32 = vec.iter().map(|x| x * x).sum();
|
|
if norm_sq > 0.0 {
|
|
let norm = norm_sq.sqrt();
|
|
for x in vec.iter_mut() {
|
|
*x /= norm;
|
|
}
|
|
} else {
|
|
vec[0] = 1.0;
|
|
}
|
|
|
|
vec
|
|
}
|
|
|
|
enum EmbeddingModel {
|
|
Candle(CandleEmbeddingModel),
|
|
Fallback,
|
|
}
|
|
|
|
impl EmbeddingModel {
|
|
fn new() -> Self {
|
|
match CandleEmbeddingModel::new() {
|
|
Ok(model) => EmbeddingModel::Candle(model),
|
|
Err(e) => {
|
|
tracing::warn!(
|
|
"Failed to initialize Candle BERT model ({e}); falling back to deterministic offline embeddings."
|
|
);
|
|
EmbeddingModel::Fallback
|
|
}
|
|
}
|
|
}
|
|
|
|
fn embed(&self, texts: &[String]) -> Result<Vec<Vec<f32>>, String> {
|
|
match self {
|
|
EmbeddingModel::Candle(model) => model.embed(texts),
|
|
EmbeddingModel::Fallback => Ok(texts.iter().map(|t| fallback_embed(t)).collect()),
|
|
}
|
|
}
|
|
}
|
|
|
|
static EMBEDDING_MODEL: OnceLock<Mutex<EmbeddingModel>> = OnceLock::new();
|
|
static INIT_MUTEX: Mutex<()> = Mutex::new(());
|
|
|
|
fn get_embedding_model() -> Result<&'static Mutex<EmbeddingModel>, String> {
|
|
if let Some(model) = EMBEDDING_MODEL.get() {
|
|
return Ok(model);
|
|
}
|
|
let _guard = INIT_MUTEX.lock().map_err(|e| e.to_string())?;
|
|
if let Some(model) = EMBEDDING_MODEL.get() {
|
|
return Ok(model);
|
|
}
|
|
let model = EmbeddingModel::new();
|
|
let _ = EMBEDDING_MODEL.set(Mutex::new(model));
|
|
Ok(EMBEDDING_MODEL.get().unwrap())
|
|
}
|
|
pub async fn generate_embedding_async(text: String) -> Result<Vec<f32>, String> {
|
|
tokio::task::spawn_blocking(move || {
|
|
let model_mutex = get_embedding_model()?;
|
|
let model = model_mutex.lock().map_err(|e| e.to_string())?;
|
|
let embeddings = model.embed(&[text])?;
|
|
let emb = embeddings
|
|
.into_iter()
|
|
.next()
|
|
.ok_or_else(|| "Embedding model returned no embeddings".to_string())?;
|
|
if emb.is_empty() {
|
|
return Err("Embedding model generated a 0-length vector".to_string());
|
|
}
|
|
Ok(emb)
|
|
})
|
|
.await
|
|
.map_err(|e| e.to_string())?
|
|
}
|
|
|
|
pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
|
|
if a.is_empty() || b.is_empty() || a.len() != b.len() {
|
|
return 0.0;
|
|
}
|
|
let mut dot_product = 0.0f32;
|
|
let mut norm_a_sq = 0.0f32;
|
|
let mut norm_b_sq = 0.0f32;
|
|
|
|
let chunks_a = a.chunks_exact(8);
|
|
let chunks_b = b.chunks_exact(8);
|
|
|
|
let remainder_a = chunks_a.remainder();
|
|
let remainder_b = chunks_b.remainder();
|
|
|
|
for (ca, cb) in chunks_a.zip(chunks_b) {
|
|
for i in 0..8 {
|
|
let x = ca[i];
|
|
let y = cb[i];
|
|
dot_product += x * y;
|
|
norm_a_sq += x * x;
|
|
norm_b_sq += y * y;
|
|
}
|
|
}
|
|
|
|
for (&x, &y) in remainder_a.iter().zip(remainder_b.iter()) {
|
|
dot_product += x * y;
|
|
norm_a_sq += x * x;
|
|
norm_b_sq += y * y;
|
|
}
|
|
|
|
// Fast path: If vectors are already normalized (Candle & fallback embeddings), skip square roots
|
|
if (norm_a_sq - 1.0).abs() < 1e-4 && (norm_b_sq - 1.0).abs() < 1e-4 {
|
|
return dot_product.clamp(-1.0, 1.0);
|
|
}
|
|
|
|
let norm_product = norm_a_sq * norm_b_sq;
|
|
if norm_product <= 0.0 {
|
|
0.0
|
|
} else {
|
|
(dot_product / norm_product.sqrt()).clamp(-1.0, 1.0)
|
|
}
|
|
}
|
|
|
|
pub async fn generate_embeddings_async(texts: Vec<String>) -> Result<Vec<Vec<f32>>, String> {
|
|
if texts.is_empty() {
|
|
return Ok(Vec::new());
|
|
}
|
|
tokio::task::spawn_blocking(move || {
|
|
let model_mutex = get_embedding_model()?;
|
|
let model = model_mutex.lock().map_err(|e| e.to_string())?;
|
|
let mut all_embeddings = Vec::with_capacity(texts.len());
|
|
|
|
let mut current_chunk = Vec::new();
|
|
let mut current_chars = 0;
|
|
const MAX_CHARS_PER_BATCH: usize = 16384;
|
|
|
|
for text in texts {
|
|
let text_len = text.len();
|
|
if !current_chunk.is_empty()
|
|
&& (current_chunk.len() >= 64 || current_chars + text_len > MAX_CHARS_PER_BATCH)
|
|
{
|
|
let chunk_vec = std::mem::take(&mut current_chunk);
|
|
let chunk_embeddings = model.embed(&chunk_vec)?;
|
|
all_embeddings.extend(chunk_embeddings);
|
|
current_chars = 0;
|
|
}
|
|
current_chars += text_len;
|
|
current_chunk.push(text);
|
|
}
|
|
if !current_chunk.is_empty() {
|
|
let chunk_embeddings = model.embed(¤t_chunk)?;
|
|
all_embeddings.extend(chunk_embeddings);
|
|
}
|
|
|
|
Ok(all_embeddings)
|
|
})
|
|
.await
|
|
.map_err(|e| e.to_string())?
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
|
|
#[test]
|
|
fn test_cosine_similarity_edge_cases() {
|
|
assert_eq!(cosine_similarity(&[], &[]), 0.0);
|
|
assert_eq!(cosine_similarity(&[1.0, 2.0], &[1.0]), 0.0);
|
|
assert_eq!(cosine_similarity(&[0.0, 0.0], &[0.0, 0.0]), 0.0);
|
|
|
|
let v1 = vec![1.0, 0.0, 0.0];
|
|
let v2 = vec![1.0, 0.0, 0.0];
|
|
assert!((cosine_similarity(&v1, &v2) - 1.0).abs() < 1e-5);
|
|
|
|
let v3 = vec![0.0, 1.0, 0.0];
|
|
assert!((cosine_similarity(&v1, &v3) - 0.0).abs() < 1e-5);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_generate_embeddings_async_empty() {
|
|
let res = generate_embeddings_async(vec![]).await.unwrap();
|
|
assert!(res.is_empty());
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_generate_embeddings_async_single_text() {
|
|
let text = "test text".to_string();
|
|
let res = generate_embeddings_async(vec![text.clone()]).await.unwrap();
|
|
assert_eq!(res.len(), 1);
|
|
assert_eq!(res[0].len(), 384);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_generate_embeddings_async_multiple_texts() {
|
|
let texts = vec![
|
|
"test text 1".to_string(),
|
|
"test text 2".to_string(),
|
|
"test text 3".to_string(),
|
|
];
|
|
let res = generate_embeddings_async(texts.clone()).await.unwrap();
|
|
assert_eq!(res.len(), 3);
|
|
for embedding in &res {
|
|
assert_eq!(embedding.len(), 384);
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_fallback_embed_deterministic_stability() {
|
|
let text = "The quick brown fox jumps over the lazy dog";
|
|
let emb1 = fallback_embed(text);
|
|
let emb2 = fallback_embed(text);
|
|
assert_eq!(emb1.len(), 384);
|
|
assert_eq!(emb1, emb2);
|
|
assert!((cosine_similarity(&emb1, &emb2) - 1.0).abs() < 1e-5);
|
|
}
|
|
}
|