From 84e9515509c559ed01e4b0a67539f10cd2c065e6 Mon Sep 17 00:00:00 2001 From: sigoden Date: Mon, 9 Sep 2024 20:26:59 +0800 Subject: refactor: abandon config `rag_min_score_rerank` (#852) --- src/config/mod.rs | 22 ++++------------------ src/rag/mod.rs | 40 +++++++++++++++++++++++----------------- 2 files changed, 27 insertions(+), 35 deletions(-) (limited to 'src') diff --git a/src/config/mod.rs b/src/config/mod.rs index 25f0cd6..5cd266f 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -9,8 +9,8 @@ pub use self::role::{Role, RoleLike, BUILTIN_ROLES, CODE_ROLE, EXPLAIN_SHELL_ROL use self::session::Session; use crate::client::{ - create_client_config, init_client, list_chat_models, list_client_types, list_reranker_models, - ClientConfig, Model, OPENAI_COMPATIBLE_PLATFORMS, + create_client_config, list_chat_models, list_client_types, list_reranker_models, ClientConfig, + Model, OPENAI_COMPATIBLE_PLATFORMS, }; use crate::function::{FunctionDeclaration, Functions, ToolResult}; use crate::rag::Rag; @@ -117,7 +117,6 @@ pub struct Config { pub rag_chunk_overlap: Option, pub rag_min_score_vector_search: f32, pub rag_min_score_keyword_search: f32, - pub rag_min_score_rerank: f32, pub rag_template: Option, #[serde(default)] @@ -185,7 +184,6 @@ impl Default for Config { rag_chunk_overlap: None, rag_min_score_vector_search: 0.0, rag_min_score_keyword_search: 0.0, - rag_min_score_rerank: 0.0, rag_template: None, document_loaders: Default::default(), @@ -1146,29 +1144,20 @@ impl Config { abort_signal: AbortSignal, ) -> Result { let (reranker_model, top_k) = rag.get_config(); - let (min_score_vector_search, min_score_keyword_search, rag_min_score_rerank) = { + let (min_score_vector_search, min_score_keyword_search) = { let config = config.read(); ( config.rag_min_score_vector_search, config.rag_min_score_keyword_search, - config.rag_min_score_rerank, ) }; - let rerank = match reranker_model { - Some(reranker_model_id) => { - let rerank_model = Model::retrieve_reranker(&config.read(), &reranker_model_id)?; - let rerank_client = init_client(config, Some(rerank_model))?; - Some((rerank_client, rag_min_score_rerank)) - } - None => None, - }; let embeddings = rag .search( text, top_k, min_score_vector_search, min_score_keyword_search, - rerank, + reranker_model.as_deref(), abort_signal, ) .await?; @@ -1849,9 +1838,6 @@ impl Config { if let Some(Some(v)) = read_env_value::("rag_min_score_keyword_search") { self.rag_min_score_keyword_search = v; } - if let Some(Some(v)) = read_env_value::("rag_min_score_rerank") { - self.rag_min_score_rerank = v; - } if let Some(v) = read_env_value::("rag_template") { self.rag_template = v; } diff --git a/src/rag/mod.rs b/src/rag/mod.rs index e03b3f1..42e6582 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -285,12 +285,12 @@ impl Rag { top_k: usize, min_score_vector_search: f32, min_score_keyword_search: f32, - rerank: Option<(Box, f32)>, + rerank_model: Option<&str>, abort_signal: AbortSignal, ) -> Result { let spinner = create_spinner("Searching").await; let ret = tokio::select! { - ret = self.hybird_search(text, top_k, min_score_vector_search, min_score_keyword_search, rerank) => { + ret = self.hybird_search(text, top_k, min_score_vector_search, min_score_keyword_search, rerank_model) => { ret } _ = watch_abort_signal(abort_signal) => { @@ -425,7 +425,7 @@ impl Rag { top_k: usize, min_score_vector_search: f32, min_score_keyword_search: f32, - rerank: Option<(Box, f32)>, + rerank_model: Option<&str>, ) -> Result> { let (vector_search_result, text_search_result) = tokio::join!( self.vector_search(query, top_k, min_score_vector_search), @@ -434,11 +434,14 @@ impl Rag { let vector_search_ids = vector_search_result?; let keyword_search_ids = text_search_result?; debug!( - "vector_search_ids: {vector_search_ids:?}, keyword_search_ids: {keyword_search_ids:?}" + "vector_search_ids: {:?}, keyword_search_ids: {:?}", + pretty_document_ids(&vector_search_ids), + pretty_document_ids(&keyword_search_ids) ); - let ids = match rerank { - Some((client, min_score)) => { - let min_score = min_score as f64; + let ids = match rerank_model { + Some(model_id) => { + let model = Model::retrieve_reranker(&self.config.read(), model_id)?; + let client = init_client(&self.config, Some(model))?; let ids: IndexSet = [vector_search_ids, keyword_search_ids] .concat() .into_iter() @@ -453,18 +456,12 @@ impl Rag { } let data = RerankData::new(query.to_string(), documents, top_k); let list = client.rerank(data).await?; - let ids = list + let ids: Vec<_> = list .into_iter() .take(top_k) - .filter_map(|item| { - if item.relevance_score < min_score { - None - } else { - documents_ids.get(item.index).cloned() - } - }) + .filter_map(|item| documents_ids.get(item.index).cloned()) .collect(); - debug!("rerank_ids: {ids:?}"); + debug!("rerank_ids: {:?}", pretty_document_ids(&ids)); ids } None => { @@ -473,7 +470,7 @@ impl Rag { vec![1.0, 1.0], top_k, ); - debug!("rrf_ids: {ids:?}"); + debug!("rrf_ids: {:?}", pretty_document_ids(&ids)); ids } }; @@ -713,6 +710,15 @@ pub fn split_document_id(value: DocumentId) -> (usize, usize) { (high, low) } +fn pretty_document_ids(ids: &[DocumentId]) -> Vec { + ids.iter() + .map(|v| { + let (h, l) = split_document_id(*v); + format!("{h}-{l}") + }) + .collect() +} + fn select_embedding_model(models: &[&Model]) -> Result { let models: Vec<_> = models .iter() -- cgit v1.2.3