diff --git a/embeddings/src/model/ffi_test.rs b/embeddings/src/model/ffi_test.rs index bc33c358..8e4c4dde 100644 --- a/embeddings/src/model/ffi_test.rs +++ b/embeddings/src/model/ffi_test.rs @@ -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 = (0..500) + .map(|i| format!("large concurrent BERT batch {worker} document {i}")) + .collect(); + let items: Vec = + 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. diff --git a/embeddings/src/model/local.rs b/embeddings/src/model/local.rs index 67bab744..571066db 100644 --- a/embeddings/src/model/local.rs +++ b/embeddings/src/model/local.rs @@ -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]) -> Result>, Box> { + fn predict_chunks( + &self, + model: &BertModel, + chunks: &[Vec], + ) -> Result>, Box> { let mut all_embeddings = Vec::with_capacity(chunks.len()); for batch in chunks.chunks(batch_size()) { @@ -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)?; @@ -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) @@ -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>, Box> { + fn predict_local( + &self, + mut locked: Locked<'_>, + texts: &[&str], + ) -> Result>, Box> { // 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. @@ -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)?; @@ -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) }); } @@ -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, @@ -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!(), }; @@ -1497,14 +1507,38 @@ impl LocalModel { impl TextModel for LocalModel { fn predict(&self, texts: &[&str], threads: usize) -> Result>, Box> { - // 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 {