diff options
Diffstat (limited to 'src/config')
| -rw-r--r-- | src/config/input.rs | 21 | ||||
| -rw-r--r-- | src/config/mod.rs | 29 |
2 files changed, 39 insertions, 11 deletions
diff --git a/src/config/input.rs b/src/config/input.rs index 892a41f..03ec999 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -169,20 +169,31 @@ impl Input { if !self.text.is_empty() { let rag = self.config.read().rag.clone(); if let Some(rag) = rag { - let (top_k, min_score_vector, min_score_text) = { + let (top_k, min_score_vector_search, min_score_fulltext_search) = { let config = self.config.read(); ( config.rag_top_k, - config.rag_min_score_vector, - config.rag_min_score_text, + config.rag_min_score_vector_search, + config.rag_min_score_fulltext_search, ) }; + let rerank = match self.config.read().rag_rerank_model.clone() { + Some(rerank_model_id) => { + let min_score = self.config.read().rag_min_score_rerank; + let rerank_model = + Model::retrieve_rerank(&self.config.read(), &rerank_model_id)?; + let rerank_client = init_client(&self.config, Some(rerank_model))?; + Some((rerank_client, min_score)) + } + None => None, + }; let embeddings = rag .search( &self.text, top_k, - min_score_vector, - min_score_text, + min_score_vector_search, + min_score_fulltext_search, + rerank, abort_signal, ) .await?; diff --git a/src/config/mod.rs b/src/config/mod.rs index 509e4de..e07dca6 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -9,8 +9,8 @@ pub use self::role::{Role, RoleLike, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE}; use self::session::Session; use crate::client::{ - create_client_config, list_chat_models, list_client_types, ClientConfig, Model, - OPENAI_COMPATIBLE_PLATFORMS, + create_client_config, list_chat_models, list_client_types, list_rerank_models, ClientConfig, + Model, OPENAI_COMPATIBLE_PLATFORMS, }; use crate::function::{FunctionDeclaration, Functions, FunctionsFilter, ToolResult}; use crate::rag::Rag; @@ -102,11 +102,13 @@ pub struct Config { pub bot_prelude: Option<String>, pub bots: Vec<BotConfig>, pub rag_embedding_model: Option<String>, + pub rag_rerank_model: Option<String>, pub rag_chunk_size: Option<usize>, pub rag_chunk_overlap: Option<usize>, pub rag_top_k: usize, - pub rag_min_score_vector: f32, - pub rag_min_score_text: f32, + pub rag_min_score_vector_search: f32, + pub rag_min_score_fulltext_search: f32, + pub rag_min_score_rerank: f32, pub rag_template: Option<String>, pub compress_threshold: usize, pub summarize_prompt: Option<String>, @@ -156,11 +158,13 @@ impl Default for Config { bot_prelude: None, bots: vec![], rag_embedding_model: None, + rag_rerank_model: None, rag_chunk_size: None, rag_chunk_overlap: None, rag_top_k: 4, - rag_min_score_vector: 0.0, - rag_min_score_text: 0.0, + rag_min_score_vector_search: 0.0, + rag_min_score_fulltext_search: 0.0, + rag_min_score_rerank: 0.0, rag_template: None, compress_threshold: 4000, summarize_prompt: None, @@ -442,6 +446,10 @@ impl Config { ), ("temperature", format_option_value(&role.temperature())), ("top_p", format_option_value(&role.top_p())), + ( + "rag_rerank_model", + format_option_value(&self.rag_rerank_model), + ), ("rag_top_k", self.rag_top_k.to_string()), ("function_calling", self.function_calling.to_string()), ("compress_threshold", self.compress_threshold.to_string()), @@ -490,6 +498,13 @@ impl Config { let value = parse_value(value)?; self.set_top_p(value); } + "rag_rerank_model" => { + self.rag_rerank_model = if value == "null" { + None + } else { + Some(value.to_string()) + } + } "rag_top_k" => { if let Some(value) = parse_value(value)? { self.rag_top_k = value; @@ -1052,6 +1067,7 @@ impl Config { "max_output_tokens", "temperature", "top_p", + "rag_rerank_model", "rag_top_k", "function_calling", "compress_threshold", @@ -1072,6 +1088,7 @@ impl Config { Some(v) => vec![v.to_string()], None => vec![], }, + "rag_rerank_model" => list_rerank_models(self).iter().map(|v| v.id()).collect(), "function_calling" => complete_bool(self.function_calling), "save" => complete_bool(self.save), "save_session" => { |
