From 476d29c40aaf2b7a49e2ba5068a12f789defcb75 Mon Sep 17 00:00:00 2001 From: sigoden Date: Tue, 3 Sep 2024 12:22:55 +0800 Subject: feat: use dynamic batch size for embedding (#826) --- src/rag/mod.rs | 16 +++++++++++++--- 1 file changed, 13 insertions(+), 3 deletions(-) (limited to 'src/rag/mod.rs') diff --git a/src/rag/mod.rs b/src/rag/mod.rs index 6dbb6be..1b72867 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -498,8 +498,18 @@ impl Rag { spinner: Option, ) -> Result { let EmbeddingsData { texts, query } = data; + let size = match self.embedding_model.max_input_tokens() { + Some(max_input_tokens) => { + let x = max_input_tokens / self.data.chunk_size; + match self.embedding_model.max_batch_size() { + Some(y) => x.min(y), + None => x, + } + } + None => self.embedding_model.max_batch_size().unwrap_or(1), + }; let mut output = vec![]; - let batch_chunks = texts.chunks(self.embedding_model.max_batch_size()); + let batch_chunks = texts.chunks(size.max(1)); let batch_chunks_len = batch_chunks.len(); for (index, texts) in batch_chunks.enumerate() { progress( @@ -667,8 +677,8 @@ fn select_embedding_model(models: &[&Model]) -> Result { fn set_chunk_size(model: &Model) -> Result { let default_value = model.default_chunk_size().to_string(); let help_message = model - .max_input_tokens() - .map(|v| format!("The model's max_input_token is {v}")); + .max_tokens_per_chunk() + .map(|v| format!("The model's max_tokens is {v}")); let mut text = Text::new("Set chunk size:") .with_default(&default_value) -- cgit v1.2.3