Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
70 changes: 70 additions & 0 deletions embeddings/src/model/ffi_test.rs
Original file line number Diff line number Diff line change
Expand Up @@ -525,6 +525,76 @@ mod tests {
run_concurrent_ffi_embeddings("Qwen/Qwen3-Embedding-0.6B");
}

#[test]
fn test_concurrent_large_bert_batches_via_ffi() {
use std::sync::{Arc, Barrier};
use std::thread;

let model_name = to_c_string("sentence-transformers/all-MiniLM-L6-v2");
let empty = to_c_string("");
let loaded = TextModelWrapper::load_model(
model_name.as_ptr(),
model_name.as_bytes().len(),
empty.as_ptr(),
0,
empty.as_ptr(),
0,
empty.as_ptr(),
0,
0,
false,
0,
);
if loaded.model.is_null() {
TextModelWrapper::free_model_result(loaded);
eprintln!("skipping: all-MiniLM-L6-v2 not available locally");
return;
}

let start = Arc::new(Barrier::new(3));
let model_ptr = loaded.model as usize;
let handles: Vec<_> = (0..2)
.map(|worker| {
let start = Arc::clone(&start);
thread::spawn(move || {
let texts: Vec<String> = (0..500)
.map(|i| format!("large concurrent BERT batch {worker} document {i}"))
.collect();
let items: Vec<StringItem> =
texts.iter().map(|text| create_string_item(text)).collect();
let wrapper = unsafe {
std::mem::transmute::<*mut std::ffi::c_void, TextModelWrapper>(
model_ptr as *mut std::ffi::c_void,
)
};
start.wait();
// The deadlock needs the second request to arrive while the
// first forward is already mid-flight in the pool; requests
// that start simultaneously merely serialize.
if worker == 1 {
thread::sleep(std::time::Duration::from_millis(300));
}
let result = TextModelWrapper::make_vect_embeddings(
&wrapper,
items.as_ptr(),
items.len(),
std::ptr::null(),
2,
);
assert!(result.error.is_null());
assert_eq!(result.len, items.len());
TextModelWrapper::free_vec_result(result);
})
})
.collect();

start.wait();
for handle in handles {
handle.join().unwrap();
}
TextModelWrapper::free_model_result(loaded);
}

/// End-to-end on the cached MiniLM model: all five strategies. truncate/mean
/// return one vector/doc; fixed/recursive/sentence return N vectors/doc
/// grouped by row_offsets. Verifies offsets, normalization, and clean frees.
Expand Down
122 changes: 78 additions & 44 deletions embeddings/src/model/local.rs
Original file line number Diff line number Diff line change
Expand Up @@ -548,7 +548,11 @@ impl BertEmbeddingModel {

/// Batched forward pass for multiple token sequences.
/// Groups chunks into batches, pads to uniform length, runs one forward per batch.
fn predict_chunks(&self, chunks: &[Vec<u32>]) -> Result<Vec<Vec<f32>>, Box<dyn Error>> {
fn predict_chunks(
&self,
model: &BertModel,
chunks: &[Vec<u32>],
) -> Result<Vec<Vec<f32>>, Box<dyn Error>> {
let mut all_embeddings = Vec::with_capacity(chunks.len());

for batch in chunks.chunks(batch_size()) {
Expand All @@ -559,10 +563,7 @@ impl BertEmbeddingModel {
let chunk = &batch[0];
let token_ids = Tensor::new(chunk.as_slice(), &self.device)?.unsqueeze(0)?;
let token_type_ids = token_ids.zeros_like()?;
let emb = {
let model = self.model.lock().unwrap();
model.forward(&token_ids, &token_type_ids, None)?
};
let emb = model.forward(&token_ids, &token_type_ids, None)?;
let seq_len = token_ids.dims()[1];
let summed = emb.sum(1)?.to_dtype(DType::F32)?;
let divisor = Tensor::new(seq_len as f32, &self.device)?;
Expand Down Expand Up @@ -592,10 +593,7 @@ impl BertEmbeddingModel {
Tensor::from_vec(flat_mask.clone(), (batch_size, max_len), &self.device)?;
let token_type_ids = token_ids.zeros_like()?;

let emb = {
let model = self.model.lock().unwrap();
model.forward(&token_ids, &token_type_ids, Some(&attention_mask))?
};
let emb = model.forward(&token_ids, &token_type_ids, Some(&attention_mask))?;
// emb: [batch_size, max_len, hidden_size]

// Attention-mask-aware mean pooling: sum(emb * mask) / sum(mask)
Expand Down Expand Up @@ -1316,12 +1314,32 @@ impl LocalModel {
}
}

/// Mutex-guarded model state, locked by `predict` on the caller thread before the
/// rayon pool is entered. Locking inside the pool can deadlock it: the worker that
/// parks on the mutex may be holding a stolen piece of the lock owner's forward,
/// which then never completes (manticoresearch#4915). Causal models keep no
/// shared mutable state.
enum Locked<'a> {
Bert(&'a BertModel),
T5(&'a mut T5EncoderModel),
Causal,
QuantizedGemma(&'a mut QuantizedGemmaModel),
QuantizedLlama(&'a mut QuantizedLlamaModel),
}

impl LocalModel {
/// Inner predict body for non-ONNX local models (BERT / T5 / Causal / Quantized).
/// Pulled out of the trait impl so the caller can wrap it in a scoped rayon pool.
fn predict_local(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, Box<dyn Error>> {
fn predict_local(
&self,
mut locked: Locked<'_>,
texts: &[&str],
) -> Result<Vec<Vec<f32>>, Box<dyn Error>> {
// BERT: batched path (batch_size up to batch_size() per forward pass)
if let LocalModel::Bert(m) = self {
let Locked::Bert(model) = locked else {
unreachable!()
};
// Dedicated single-text bypass: SELECT KNN(field, k, 'text') hits this
// path on every query. Skip all batching wrappers, intermediate Vecs,
// and the chunks.chunks() loop — go straight encode → forward → pool.
Expand All @@ -1336,10 +1354,7 @@ impl LocalModel {

let token_ids = Tensor::new(ids, &m.device)?.unsqueeze(0)?;
let token_type_ids = token_ids.zeros_like()?;
let emb = {
let model = m.model.lock().unwrap();
model.forward(&token_ids, &token_type_ids, None)?
};
let emb = model.forward(&token_ids, &token_type_ids, None)?;
let seq_len = token_ids.dims()[1];
let summed = emb.sum(1)?.to_dtype(DType::F32)?;
let divisor = Tensor::new(seq_len as f32, &m.device)?;
Expand All @@ -1350,7 +1365,7 @@ impl LocalModel {
}

return Self::predict_batched(&m.tokenizer, m.max_input_len, texts, |chunks| {
m.predict_chunks(chunks)
m.predict_chunks(model, chunks)
});
}

Expand Down Expand Up @@ -1393,15 +1408,14 @@ impl LocalModel {

{
let token_ids = Tensor::new(&tokens[..], &device)?.unsqueeze(0)?;
let embeddings = match self {
LocalModel::T5(m) => {
let mut model = m.model.lock().unwrap();
let embeddings = match (self, &mut locked) {
(_, Locked::T5(model)) => {
let emb = model.forward(&token_ids)?;
let cls_emb = emb.i(0)?;
let first_token = cls_emb.i(0)?;
first_token.unsqueeze(0)?.to_dtype(DType::F32)?
}
LocalModel::Causal(m) => match &m.kind {
(LocalModel::Causal(m), Locked::Causal) => match &m.kind {
CausalEmbeddingKind::Qwen { model, config } => qwen_mean_pool(
model,
config,
Expand Down Expand Up @@ -1440,24 +1454,20 @@ impl LocalModel {
summed.broadcast_div(&divisor)?
}
},
LocalModel::Quantized(m) => match &m.model {
QuantizedModelKind::Gemma { model } => {
let mut model = model.lock().unwrap();
let emb = model.forward(&token_ids, 0)?;
let (_, n_tokens, _) = emb.dims3()?;
let summed = emb.sum(1)?.to_dtype(DType::F32)?;
let divisor = Tensor::new(n_tokens as f32, &device)?;
summed.broadcast_div(&divisor)?
}
QuantizedModelKind::Llama { model } => {
let mut model = model.lock().unwrap();
let emb = model.forward(&token_ids, 0)?;
let (_, n_tokens, _) = emb.dims3()?;
let summed = emb.sum(1)?.to_dtype(DType::F32)?;
let divisor = Tensor::new(n_tokens as f32, &device)?;
summed.broadcast_div(&divisor)?
}
},
(_, Locked::QuantizedGemma(model)) => {
let emb = model.forward(&token_ids, 0)?;
let (_, n_tokens, _) = emb.dims3()?;
let summed = emb.sum(1)?.to_dtype(DType::F32)?;
let divisor = Tensor::new(n_tokens as f32, &device)?;
summed.broadcast_div(&divisor)?
}
(_, Locked::QuantizedLlama(model)) => {
let emb = model.forward(&token_ids, 0)?;
let (_, n_tokens, _) = emb.dims3()?;
let summed = emb.sum(1)?.to_dtype(DType::F32)?;
let divisor = Tensor::new(n_tokens as f32, &device)?;
summed.broadcast_div(&divisor)?
}
_ => unreachable!(),
};

Expand Down Expand Up @@ -1497,14 +1507,38 @@ impl LocalModel {

impl TextModel for LocalModel {
fn predict(&self, texts: &[&str], threads: usize) -> Result<Vec<Vec<f32>>, Box<dyn Error>> {
// ONNX manages its own worker count internally — no rayon pool involved.
if let LocalModel::Onnx(m) = self {
return m.predict_pipelined(texts, threads);
}

// BERT / T5 / Causal / Quantized go through candle, which uses rayon for
// intra-op parallelism. Scope the rayon pool so threads > 0 caps the worker count.
with_thread_limit(threads, || self.predict_local(texts))
// intra-op parallelism. Scope the rayon pool so threads > 0 caps the worker
// count. Model mutexes are taken here, on the caller thread — see `Locked`.
match self {
// ONNX manages its own worker count internally — no rayon pool involved.
LocalModel::Onnx(m) => m.predict_pipelined(texts, threads),
LocalModel::Bert(m) => {
let model = m.model.lock().unwrap();
let locked = Locked::Bert(&model);
with_thread_limit(threads, move || self.predict_local(locked, texts))
}
LocalModel::T5(m) => {
let mut model = m.model.lock().unwrap();
let locked = Locked::T5(&mut model);
with_thread_limit(threads, move || self.predict_local(locked, texts))
}
LocalModel::Causal(_) => {
with_thread_limit(threads, || self.predict_local(Locked::Causal, texts))
}
LocalModel::Quantized(m) => match &m.model {
QuantizedModelKind::Gemma { model } => {
let mut model = model.lock().unwrap();
let locked = Locked::QuantizedGemma(&mut model);
with_thread_limit(threads, move || self.predict_local(locked, texts))
}
QuantizedModelKind::Llama { model } => {
let mut model = model.lock().unwrap();
let locked = Locked::QuantizedLlama(&mut model);
with_thread_limit(threads, move || self.predict_local(locked, texts))
}
},
}
}

fn get_hidden_size(&self) -> usize {
Expand Down