summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/client/model.rs17
-rw-r--r--src/rag/mod.rs16
2 files changed, 24 insertions, 9 deletions
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<usize> {
+ 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<usize> {
+ self.data.max_batch_size
}
pub fn max_tokens_param(&self) -> Option<isize> {
@@ -266,6 +270,7 @@ pub struct ModelData {
pub supports_function_calling: bool,
// embedding-only properties
+ pub max_tokens_per_chunk: Option<usize>,
pub default_chunk_size: Option<usize>,
pub max_batch_size: Option<usize>,
}
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)