diff options
| author | sigoden <sigoden@gmail.com> | 2024-09-03 12:22:55 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-09-03 12:22:55 +0800 |
| commit | 476d29c40aaf2b7a49e2ba5068a12f789defcb75 (patch) | |
| tree | 85dd550ddce734c173cb28cb6b2e8e8f60cb0190 /src/rag | |
| parent | 3695c12646fccb0e07b6dee4cdf54dfaa8159229 (diff) | |
| download | aichat-476d29c40aaf2b7a49e2ba5068a12f789defcb75.tar.gz | |
feat: use dynamic batch size for embedding (#826)
Diffstat (limited to 'src/rag')
| -rw-r--r-- | src/rag/mod.rs | 16 |
1 files changed, 13 insertions, 3 deletions
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<Spinner>, ) -> Result<EmbeddingsOutput> { 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<String> { fn set_chunk_size(model: &Model) -> Result<usize> { 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) |
