diff options
| author | sigoden <sigoden@gmail.com> | 2024-06-21 21:56:25 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-06-21 21:56:25 +0800 |
| commit | 7a089d846ec2b0f5447d253a52e884a5a6607881 (patch) | |
| tree | e6740bc652c2b93e6f6a4bfa8ace722da8d1f2cf /src | |
| parent | f2378e172548f3357fca6e561803bfe3a013bcb3 (diff) | |
| download | aichat-7a089d846ec2b0f5447d253a52e884a5a6607881.tar.gz | |
refactor: rename model.max_concurrent_chunks to model.max_batch_size (#626)
Diffstat (limited to 'src')
| -rw-r--r-- | src/client/common.rs | 2 | ||||
| -rw-r--r-- | src/client/model.rs | 12 | ||||
| -rw-r--r-- | src/rag/mod.rs | 10 |
3 files changed, 12 insertions, 12 deletions
diff --git a/src/client/common.rs b/src/client/common.rs index de2fa50..796c632 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -392,7 +392,7 @@ pub trait Client: Sync + Send { async fn embeddings(&self, data: EmbeddingsData) -> Result<Vec<Vec<f32>>> { let client = self.build_client()?; - self.model().guard_max_concurrent_chunks(&data)?; + self.model().guard_max_batch_size(&data)?; self.embeddings_inner(&client, data) .await .context("Failed to fetch embeddings") diff --git a/src/client/model.rs b/src/client/model.rs index 0584f82..4c69d78 100644 --- a/src/client/model.rs +++ b/src/client/model.rs @@ -175,8 +175,8 @@ impl Model { self.data.default_chunk_size.unwrap_or(1000) } - pub fn max_concurrent_chunks(&self) -> usize { - self.data.max_concurrent_chunks.unwrap_or(1) + pub fn max_batch_size(&self) -> usize { + self.data.max_batch_size.unwrap_or(1) } pub fn max_tokens_param(&self) -> Option<isize> { @@ -234,9 +234,9 @@ impl Model { Ok(()) } - pub fn guard_max_concurrent_chunks(&self, data: &EmbeddingsData) -> Result<()> { - if data.texts.len() > self.max_concurrent_chunks() { - bail!("Exceed max_concurrent_chunks limit"); + pub fn guard_max_batch_size(&self, data: &EmbeddingsData) -> Result<()> { + if data.texts.len() > self.max_batch_size() { + bail!("Exceed max_batch_size limit"); } Ok(()) } @@ -262,7 +262,7 @@ pub struct ModelData { // embedding-only properties pub default_chunk_size: Option<usize>, - pub max_concurrent_chunks: Option<usize>, + pub max_batch_size: Option<usize>, } impl ModelData { diff --git a/src/rag/mod.rs b/src/rag/mod.rs index 628a9e4..116beea 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -414,13 +414,13 @@ impl Rag { ) -> Result<EmbeddingsOutput> { let EmbeddingsData { texts, query } = data; let mut output = vec![]; - let chunks = texts.chunks(self.embedding_model.max_concurrent_chunks()); - let chunks_len = chunks.len(); + let batch_chunks = texts.chunks(self.embedding_model.max_batch_size()); + let batch_chunks_len = batch_chunks.len(); progress( &progress_tx, - format!("Creating embeddings [1/{chunks_len}]"), + format!("Creating embeddings [1/{batch_chunks_len}]"), ); - for (index, texts) in chunks.enumerate() { + for (index, texts) in batch_chunks.enumerate() { let chunk_data = EmbeddingsData { texts: texts.to_vec(), query, @@ -433,7 +433,7 @@ impl Rag { output.extend(chunk_output); progress( &progress_tx, - format!("Creating embeddings [{}/{chunks_len}]", index + 1), + format!("Creating embeddings [{}/{batch_chunks_len}]", index + 1), ); } Ok(output) |
