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/client/model.rs | 17 +++++++++++------ src/rag/mod.rs | 16 +++++++++++++--- 2 files changed, 24 insertions(+), 9 deletions(-) (limited to 'src') diff --git a/src/client/model.rs b/src/client/model.rs index 50dabeb..08ced10 100644 --- a/src/client/model.rs +++ b/src/client/model.rs @@ -157,15 +157,15 @@ impl Model { } "embedding" => { let ModelData { - max_input_tokens, input_price, + max_tokens_per_chunk, max_batch_size, .. } = &self.data; - let max_tokens = format_option_value(max_input_tokens); + let max_tokens = format_option_value(max_tokens_per_chunk); + let max_batch = format_option_value(max_batch_size); let price = format_option_value(input_price); - let batch = format_option_value(max_batch_size); - format!("max-tokens:{max_tokens}; price:{price}; batch:{batch}") + format!("max-tokens:{max_tokens};max-batch:{max_batch};price:{price}") } _ => String::new(), } @@ -183,12 +183,16 @@ impl Model { self.data.supports_vision } + pub fn max_tokens_per_chunk(&self) -> Option { + self.data.max_tokens_per_chunk + } + pub fn default_chunk_size(&self) -> usize { self.data.default_chunk_size.unwrap_or(1000) } - pub fn max_batch_size(&self) -> usize { - self.data.max_batch_size.unwrap_or(1) + pub fn max_batch_size(&self) -> Option { + self.data.max_batch_size } pub fn max_tokens_param(&self) -> Option { @@ -266,6 +270,7 @@ pub struct ModelData { pub supports_function_calling: bool, // embedding-only properties + pub max_tokens_per_chunk: Option, pub default_chunk_size: Option, pub max_batch_size: Option, } 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