From 89554e0d4ebcc66442456f07403be2a70a50386b Mon Sep 17 00:00:00 2001 From: sigoden Date: Sun, 8 Sep 2024 17:19:00 +0800 Subject: feat: support RAG-scoped rag_top_k and rag_reranker_model (#847) --- src/config/mod.rs | 104 +++++++++++++++++++++++++++++++++++++----------------- src/rag/mod.rs | 86 +++++++++++++++++++++++++++++++++++--------- src/repl/mod.rs | 2 +- 3 files changed, 141 insertions(+), 51 deletions(-) (limited to 'src') diff --git a/src/config/mod.rs b/src/config/mod.rs index a5b5c5b..d2d7c2f 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -511,6 +511,10 @@ impl Config { .wrap .clone() .map_or_else(|| String::from("no"), |v| v.to_string()); + let (rag_reranker_model, rag_top_k) = match self.rag.as_ref() { + Some(rag) => rag.get_config(), + None => (self.rag_reranker_model.clone(), self.rag_top_k), + }; let role = self.extract_role(); let mut items = vec![ ("model", role.model().id()), @@ -535,9 +539,9 @@ impl Config { ("use_tools", format_option_value(&role.use_tools())), ( "rag_reranker_model", - format_option_value(&self.rag_reranker_model), + format_option_value(&rag_reranker_model), ), - ("rag_top_k", self.rag_top_k.to_string()), + ("rag_top_k", rag_top_k.to_string()), ("highlight", self.highlight.to_string()), ("light_theme", self.light_theme.to_string()), ("config_file", display_path(&Self::config_file()?)), @@ -559,7 +563,7 @@ impl Config { Ok(output) } - pub fn update(&mut self, data: &str) -> Result<()> { + pub fn update(config: &GlobalConfig, data: &str) -> Result<()> { let parts: Vec<&str> = data.split_whitespace().collect(); if parts.len() != 2 { bail!("Usage: .set . If value is null, unset key."); @@ -569,62 +573,58 @@ impl Config { match key { "max_output_tokens" => { let value = parse_value(value)?; - self.set_max_output_tokens(value); + config.write().set_max_output_tokens(value); } "temperature" => { let value = parse_value(value)?; - self.set_temperature(value); + config.write().set_temperature(value); } "top_p" => { let value = parse_value(value)?; - self.set_top_p(value); + config.write().set_top_p(value); } "dry_run" => { let value = value.parse().with_context(|| "Invalid value")?; - self.dry_run = value; + config.write().dry_run = value; } "stream" => { let value = value.parse().with_context(|| "Invalid value")?; - self.stream = value; + config.write().stream = value; } "save" => { let value = value.parse().with_context(|| "Invalid value")?; - self.save = value; + config.write().save = value; } "rag_reranker_model" => { - self.rag_reranker_model = if value == "null" { - None - } else { - Some(value.to_string()) - } + let value = parse_value(value)?; + Self::set_rag_reranker_model(config, value)?; } "rag_top_k" => { - if let Some(value) = parse_value(value)? { - self.rag_top_k = value; - } + let value = value.parse().with_context(|| "Invalid value")?; + Self::set_rag_top_k(config, value)?; } "function_calling" => { let value = value.parse().with_context(|| "Invalid value")?; - if value && self.functions.is_empty() { + if value && config.write().functions.is_empty() { bail!("Function calling cannot be enabled because no functions are installed.") } - self.function_calling = value; + config.write().function_calling = value; } "use_tools" => { let value = parse_value(value)?; - self.set_use_tools(value); + config.write().set_use_tools(value); } "save_session" => { let value = parse_value(value)?; - self.set_save_session(value); + config.write().set_save_session(value); } "compress_threshold" => { let value = parse_value(value)?; - self.set_compress_threshold(value); + config.write().set_compress_threshold(value); } "highlight" => { let value = value.parse().with_context(|| "Invalid value")?; - self.highlight = value; + config.write().highlight = value; } _ => bail!("Unknown key `{key}`"), } @@ -668,6 +668,33 @@ impl Config { } } + pub fn set_rag_reranker_model(config: &GlobalConfig, value: Option) -> Result<()> { + if let Some(id) = &value { + Model::retrieve_reranker(&config.read(), id)?; + } + let has_rag = config.read().rag.is_some(); + match has_rag { + true => update_rag(config, |rag| { + rag.set_reranker_model(value)?; + Ok(()) + })?, + false => config.write().rag_reranker_model = value, + } + Ok(()) + } + + pub fn set_rag_top_k(config: &GlobalConfig, value: usize) -> Result<()> { + let has_rag = config.read().rag.is_some(); + match has_rag { + true => update_rag(config, |rag| { + rag.set_top_k(value)?; + Ok(()) + })?, + false => config.write().rag_top_k = value, + } + Ok(()) + } + pub fn set_wrap(&mut self, value: &str) -> Result<()> { if value == "no" { self.wrap = None; @@ -1090,13 +1117,11 @@ impl Config { } pub async fn rebuild_rag(config: &GlobalConfig, abort_signal: AbortSignal) -> Result<()> { - let rag_name = match config.read().rag.clone() { - Some(v) => v.name().to_string(), + let mut rag = match config.read().rag.clone() { + Some(v) => v.as_ref().clone(), None => bail!("No RAG"), }; - let rag_path = config.read().rag_file(&rag_name)?; - let mut rag = Rag::load(config, &rag_name, &rag_path)?; - rag.rebuild(config, &rag_path, abort_signal).await?; + rag.rebuild(config, abort_signal).await?; config.write().rag = Some(Arc::new(rag)); Ok(()) } @@ -1120,20 +1145,20 @@ impl Config { text: &str, abort_signal: AbortSignal, ) -> Result { - let (top_k, min_score_vector_search, min_score_keyword_search) = { + let (reranker_model, top_k) = rag.get_config(); + let (min_score_vector_search, min_score_keyword_search, rag_min_score_rerank) = { let config = config.read(); ( - config.rag_top_k, config.rag_min_score_vector_search, config.rag_min_score_keyword_search, + config.rag_min_score_rerank, ) }; - let rerank = match config.read().rag_reranker_model.clone() { + let rerank = match reranker_model { Some(reranker_model_id) => { - let min_score = config.read().rag_min_score_rerank; let rerank_model = Model::retrieve_reranker(&config.read(), &reranker_model_id)?; let rerank_client = init_client(config, Some(rerank_model))?; - Some((rerank_client, min_score)) + Some((rerank_client, rag_min_score_rerank)) } None => None, }; @@ -2056,3 +2081,16 @@ fn complete_option_bool(value: Option) -> Vec { None => vec!["true".to_string(), "false".to_string()], } } + +fn update_rag(config: &GlobalConfig, f: F) -> Result<()> +where + F: FnOnce(&mut Rag) -> Result<()>, +{ + let mut rag = match config.read().rag.clone() { + Some(v) => v.as_ref().clone(), + None => bail!("No RAG"), + }; + f(&mut rag)?; + config.write().rag = Some(Arc::new(rag)); + Ok(()) +} diff --git a/src/rag/mod.rs b/src/rag/mod.rs index aa1911e..f6a511d 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -22,13 +22,13 @@ use std::collections::HashMap; use std::{fmt::Debug, io::BufReader, path::Path}; pub struct Rag { + config: GlobalConfig, name: String, path: String, embedding_model: Model, hnsw: Hnsw<'static, f32, DistCosine>, bm25: BM25, data: RagData, - embedding_client: Box, } impl Debug for Rag { @@ -42,6 +42,20 @@ impl Debug for Rag { } } +impl Clone for Rag { + fn clone(&self) -> Self { + Self { + config: self.config.clone(), + name: self.name.clone(), + path: self.path.clone(), + embedding_model: self.embedding_model.clone(), + hnsw: self.data.build_hnsw(), + bm25: self.bm25.clone(), + data: self.data.clone(), + } + } +} + impl Rag { pub async fn init( config: &GlobalConfig, @@ -51,8 +65,18 @@ impl Rag { abort_signal: AbortSignal, ) -> Result { debug!("init rag: {name}"); - let (embedding_model, chunk_size, chunk_overlap) = Self::config(config)?; - let data = RagData::new(embedding_model.id(), chunk_size, chunk_overlap); + let (embedding_model, chunk_size, chunk_overlap) = Self::create_config(config)?; + let (reranker_model, top_k) = { + let config = config.read(); + (config.rag_reranker_model.clone(), config.rag_top_k) + }; + let data = RagData::new( + embedding_model.id(), + chunk_size, + chunk_overlap, + reranker_model, + top_k, + ); let mut rag = Self::create(config, name, save_path, data)?; let mut paths = doc_paths.to_vec(); if paths.is_empty() { @@ -71,8 +95,7 @@ impl Rag { bail!("Aborted!") }, }; - if !rag.is_temp() { - rag.save(save_path)?; + if rag.save()? { println!("✨ Saved rag to '{}'", save_path.display()); } Ok(rag) @@ -90,15 +113,14 @@ impl Rag { let hnsw = data.build_hnsw(); let bm25 = data.build_bm25(); let embedding_model = Model::retrieve_embedding(&config.read(), &data.embedding_model)?; - let embedding_client = init_client(config, Some(embedding_model.clone()))?; let rag = Rag { + config: config.clone(), name: name.to_string(), path: path.display().to_string(), data, embedding_model, hnsw, bm25, - embedding_client, }; Ok(rag) } @@ -106,7 +128,6 @@ impl Rag { pub async fn rebuild( &mut self, config: &GlobalConfig, - save_path: &Path, abort_signal: AbortSignal, ) -> Result<()> { debug!("rebuild rag: {}", self.name); @@ -123,14 +144,13 @@ impl Rag { bail!("Aborted!") }, }; - if !self.is_temp() { - self.save(save_path)?; - println!("✨ Saved rag to '{}'", save_path.display()); + if self.save()? { + println!("✨ Saved rag to '{}'", self.path); } Ok(()) } - pub fn config(config: &GlobalConfig) -> Result<(Model, usize, usize)> { + pub fn create_config(config: &GlobalConfig) -> Result<(Model, usize, usize)> { let (embedding_model_id, chunk_size, chunk_overlap) = { let config = config.read(); ( @@ -194,12 +214,32 @@ impl Rag { Ok((embedding_model, chunk_size, chunk_overlap)) } - pub fn save(&self, path: &Path) -> Result<()> { + pub fn get_config(&self) -> (Option, usize) { + (self.data.reranker_model.clone(), self.data.top_k) + } + + pub fn set_reranker_model(&mut self, reranker_model: Option) -> Result<()> { + self.data.reranker_model = reranker_model; + self.save()?; + Ok(()) + } + + pub fn set_top_k(&mut self, top_k: usize) -> Result<()> { + self.data.top_k = top_k; + self.save()?; + Ok(()) + } + + pub fn save(&self) -> Result { + if self.is_temp() { + return Ok(false); + } + let path = Path::new(&self.path); ensure_parent_exists(path)?; let mut file = std::fs::File::create(path)?; bincode::serialize_into(&mut file, &self.data) .with_context(|| format!("Failed to save rag '{}'", self.name))?; - Ok(()) + Ok(true) } pub fn export(&self) -> Result { @@ -219,6 +259,8 @@ impl Rag { "embedding_model": self.embedding_model.id(), "chunk_size": self.data.chunk_size, "chunk_overlap": self.data.chunk_overlap, + "reranker_model": self.data.reranker_model, + "top_k": self.data.top_k, "document_paths": self.data.document_paths, "files": files, }); @@ -490,6 +532,7 @@ impl Rag { data: EmbeddingsData, spinner: Option, ) -> Result { + let embedding_client = init_client(&self.config, Some(self.embedding_model.clone()))?; let EmbeddingsData { texts, query } = data; let size = match self.embedding_model.max_input_tokens() { Some(max_input_tokens) => { @@ -513,8 +556,7 @@ impl Rag { texts: texts.to_vec(), query, }; - let chunk_output = self - .embedding_client + let chunk_output = embedding_client .embeddings(chunk_data) .await .context("Failed to create embedding")?; @@ -529,6 +571,8 @@ pub struct RagData { pub embedding_model: String, pub chunk_size: usize, pub chunk_overlap: usize, + pub reranker_model: Option, + pub top_k: usize, pub next_file_id: FileId, pub document_paths: Vec, pub files: IndexMap, @@ -549,11 +593,19 @@ impl Debug for RagData { } impl RagData { - pub fn new(embedding_model: String, chunk_size: usize, chunk_overlap: usize) -> Self { + pub fn new( + embedding_model: String, + chunk_size: usize, + chunk_overlap: usize, + reranker_model: Option, + top_k: usize, + ) -> Self { Self { embedding_model, chunk_size, chunk_overlap, + reranker_model, + top_k, next_file_id: 0, document_paths: Default::default(), files: Default::default(), diff --git a/src/repl/mod.rs b/src/repl/mod.rs index eec43e8..b6c71a9 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -386,7 +386,7 @@ impl Repl { } ".set" => match args { Some(args) => { - self.config.write().update(args)?; + Config::update(&self.config, args)?; } _ => { println!("Usage: .set ...") -- cgit v1.2.3