feat(embedding,vision): candle embeddings with offline fallback, on-demand clipboard vision capture, and concurrency audit
This commit is contained in:
1 parent
5bd8b1587a
commit
e4a0fe72df
47 files changed
+6292
-3503
No files matched your search
+276
-20
@@ -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(¤t_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 {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in new issue
Block a user