feat(embedding,vision): candle embeddings with offline fallback, on-demand clipboard vision capture, and concurrency audit

This commit is contained in:
Riz Ashraf committed 2026-10-07 06:36:09 +01:00
1 parent 5bd8b1587a
commit e4a0fe72df
47 files changed
+6292 -3503

No files matched your search

+276 -20
View File
@@ -1,29 +1,250 @@
#[allow(deprecated)]
use fastembed::{EmbeddingModel, TextEmbedding};
use std::sync::Mutex;
use std::sync::OnceLock;
static EMBEDDING_MODEL: OnceLock<Mutex<TextEmbedding>> = OnceLock::new();
use candle_core::{DType, Device, Tensor};
use candle_nn::VarBuilder;
use candle_transformers::models::bert::{BertModel, Config};
use tokenizers::Tokenizer;
#[allow(deprecated)]
pub fn get_embedding_model() -> Result<&'static Mutex<TextEmbedding>, String> {
struct CandleEmbeddingModel {
model: BertModel,
tokenizer: Tokenizer,
device: Device,
}
impl CandleEmbeddingModel {
fn new() -> Result<Self, 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;
}
use std::hash::{Hash, Hasher};
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 mut hasher = std::collections::hash_map::DefaultHasher::new();
clean.hash(&mut hasher);
let h = hasher.finish();
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 mut h2 = std::collections::hash_map::DefaultHasher::new();
window.hash(&mut h2);
let hv = h2.finish();
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 options = fastembed::InitOptions::new(EmbeddingModel::AllMiniLML6V2)
.with_show_download_progress(true);
let model = TextEmbedding::try_new(options).map_err(|e| e.to_string())?;
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 mut model = model_mutex.lock().map_err(|e| e.to_string())?;
let embeddings = model.embed(vec![text], None).map_err(|e| e.to_string())?;
Ok(embeddings.into_iter().next().unwrap_or_default())
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())?
@@ -36,11 +257,29 @@ pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
let mut dot_product = 0.0f32;
let mut norm_a_sq = 0.0f32;
let mut norm_b_sq = 0.0f32;
for (&x, &y) in a.iter().zip(b.iter()) {
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;
}
let norm_a = norm_a_sq.sqrt();
let norm_b = norm_b_sq.sqrt();
if norm_a == 0.0 || norm_b == 0.0 {
@@ -49,19 +288,38 @@ pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
dot_product / (norm_a * norm_b)
}
}
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 mut model = model_mutex.lock().map_err(|e| e.to_string())?;
let model = model_mutex.lock().map_err(|e| e.to_string())?;
let mut all_embeddings = Vec::with_capacity(texts.len());
for chunk in texts.chunks(32) {
let chunk_vec = chunk.to_vec();
let chunk_embeddings = model.embed(chunk_vec, None).map_err(|e| e.to_string())?;
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(&current_chunk)?;
all_embeddings.extend(chunk_embeddings);
}
Ok(all_embeddings)
})
.await
@@ -93,7 +351,6 @@ mod tests {
}
#[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();
@@ -115,4 +372,3 @@ mod tests {
}
}
}