summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/client/common.rs2
-rw-r--r--src/client/model.rs12
-rw-r--r--src/rag/mod.rs10
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)