From 746b087111fabc10ec3f3f3e9ef3628d1eb47fd8 Mon Sep 17 00:00:00 2001 From: sigoden Date: Fri, 14 Jun 2024 19:12:18 +0800 Subject: refactor: add/modify rag-related config (#599) --- README.md | 5 --- config.example.yaml | 20 ++++++--- models.yaml | 4 +- src/client/model.rs | 11 ++++- src/config/bot.rs | 2 +- src/config/mod.rs | 30 +++++++++----- src/config/session.rs | 2 +- src/rag/mod.rs | 110 +++++++++++++++++++++++++++++++++++++------------- 8 files changed, 130 insertions(+), 54 deletions(-) diff --git a/README.md b/README.md index 8318df8..a1adc60 100644 --- a/README.md +++ b/README.md @@ -94,11 +94,6 @@ prelude: null # Set a default role or session to start with ( # if unset fallback to $EDITOR and $VISUAL buffer_editor: null -# Specifies the embedding model to use -embedding_model: null -# Determines how many relevant documents are retrieved -rag_top_k: 4 - # Compress session when token count reaches or exceeds this threshold (must be at least 1000) compress_threshold: 4000 diff --git a/config.example.yaml b/config.example.yaml index 7c87ade..d3e1ea1 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -35,19 +35,29 @@ bots: dangerously_functions_filter: null # Specifies the embedding model to use -embedding_model: null - +rag_embedding_model: null +# Specifies the chunk size +rag_chunk_size: null +# Specifies the chunk overlap +rag_chunk_overlap: null # Determines how many relevant documents are retrieved rag_top_k: 4 # Defines the query structure using variables like __CONTEXT__ and __INPUT__ to tailor searches to specific needs rag_template: | - Answer the following question based only on the provided context: + Use the following context as your learned knowledge, inside XML tags. - __CONTEXT__ + __CONTEXT__ - Question: __INPUT__ + When answer to user: + - If you don't know, just say that you don't know. + - If you don't know when you are not sure, ask for clarification. + Avoid mentioning that you obtained the information from the context. + And answer according to the language of the user's question. + + Given the context information, answer the query. + Query: __INPUT__ # Compress session when token count reaches or exceeds this threshold (must be at least 1000) compress_threshold: 4000 diff --git a/models.yaml b/models.yaml index 3ecde52..ac20e1e 100644 --- a/models.yaml +++ b/models.yaml @@ -204,12 +204,12 @@ - name: embed-english-v3.0 mode: embedding max_input_tokens: 512 - default_chunk_size: 500 + default_chunk_size: 1000 max_concurrent_chunks: 96 - name: embed-multilingual-v3.0 mode: embedding max_input_tokens: 512 - default_chunk_size: 500 + default_chunk_size: 1000 max_concurrent_chunks: 96 - platform: perplexity diff --git a/src/client/model.rs b/src/client/model.rs index d555232..56421bf 100644 --- a/src/client/model.rs +++ b/src/client/model.rs @@ -1,5 +1,5 @@ use super::{ - list_chat_models, + list_chat_models, list_embedding_models, message::{Message, MessageContent}, EmbeddingsData, }; @@ -43,13 +43,20 @@ impl Model { .collect() } - pub fn retrieve(config: &Config, model_id: &str) -> Result { + pub fn retrieve_chat(config: &Config, model_id: &str) -> Result { match Self::find(&list_chat_models(config), model_id) { Some(v) => Ok(v), None => bail!("Invalid model '{model_id}'"), } } + pub fn retrieve_embedding(config: &Config, model_id: &str) -> Result { + match Self::find(&list_embedding_models(config), model_id) { + Some(v) => Ok(v), + None => bail!("Invalid model '{model_id}'"), + } + } + pub fn find(models: &[&Self], model_id: &str) -> Option { let mut model = None; let (client_name, model_name) = match model_id.split_once(':') { diff --git a/src/config/bot.rs b/src/config/bot.rs index e02cce7..098c8d7 100644 --- a/src/config/bot.rs +++ b/src/config/bot.rs @@ -49,7 +49,7 @@ impl Bot { let model = { let config = config.read(); match bot_config.model_id.as_ref() { - Some(model_id) => Model::retrieve(&config, model_id)?, + Some(model_id) => Model::retrieve_chat(&config, model_id)?, None => config.current_model().clone(), } }; diff --git a/src/config/mod.rs b/src/config/mod.rs index 571bf2c..a905ae8 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -61,13 +61,19 @@ const SUMMARIZE_PROMPT: &str = "Summarize the discussion briefly in 200 words or less to use as a prompt for future context."; const SUMMARY_PROMPT: &str = "This is a summary of the chat history as a recap: "; -const RAG_TEMPLATE: &str = r#"Answer the question based only on the provided context: - -__CONTEXT__ - +const RAG_TEMPLATE: &str = r#"Use the following context as your learned knowledge, inside XML tags. + + __CONTEXT__ + -Question: __INPUT__ -"#; + When answer to user: + - If you don't know, just say that you don't know. + - If you don't know when you are not sure, ask for clarification. + Avoid mentioning that you obtained the information from the context. + And answer according to the language of the user's question. + + Given the context information, answer the query. + Query: __INPUT__"#; const LEFT_PROMPT: &str = "{color.green}{?session {?bot {bot}#}{session}{?role /}}{!session {?bot {bot}}}{role}{?rag @{rag}}{color.cyan}{?session )}{!session >}{color.reset} "; const RIGHT_PROMPT: &str = "{color.purple}{?session {?consume_tokens {consume_tokens}({consume_percent}%)}{!consume_tokens {consume_tokens}}}{color.reset}"; @@ -96,7 +102,9 @@ pub struct Config { pub dangerously_functions_filter: Option, pub bot_prelude: Option, pub bots: Vec, - pub embedding_model: Option, + pub rag_embedding_model: Option, + pub rag_chunk_size: Option, + pub rag_chunk_overlap: Option, pub rag_top_k: usize, pub rag_template: Option, pub compress_threshold: usize, @@ -147,7 +155,9 @@ impl Default for Config { dangerously_functions_filter: None, bot_prelude: None, bots: vec![], - embedding_model: None, + rag_embedding_model: None, + rag_chunk_size: None, + rag_chunk_overlap: None, rag_top_k: 4, rag_template: None, compress_threshold: 4000, @@ -616,7 +626,7 @@ impl Config { } pub fn set_model(&mut self, model_id: &str) -> Result<()> { - let model = Model::retrieve(self, model_id)?; + let model = Model::retrieve_chat(self, model_id)?; match self.role_like_mut() { Some(role_like) => role_like.set_model(&model), None => { @@ -682,7 +692,7 @@ impl Config { match role.model_id() { Some(model_id) => { if self.model.id() != model_id { - let model = Model::retrieve(self, model_id)?; + let model = Model::retrieve_chat(self, model_id)?; role.set_model(&model); } } diff --git a/src/config/session.rs b/src/config/session.rs index 38855cc..7400f30 100644 --- a/src/config/session.rs +++ b/src/config/session.rs @@ -73,7 +73,7 @@ impl Session { let mut session: Self = serde_yaml::from_str(&content).with_context(|| format!("Invalid session {}", name))?; - session.model = Model::retrieve(config, &session.model_id)?; + session.model = Model::retrieve_chat(config, &session.model_id)?; session.name = name.to_string(); session.path = Some(path.display().to_string()); diff --git a/src/rag/mod.rs b/src/rag/mod.rs index b968232..b7ac6cd 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -50,13 +50,8 @@ impl Rag { doc_paths: &[String], abort_signal: AbortSignal, ) -> Result { - if !*IS_STDOUT_TERMINAL { - bail!("An interactive shell is required to initialize rag.") - } debug!("init rag: {name}"); - let model = select_embedding_model(config)?; - let chunk_size = set_chunk_size(&model)?; - let chunk_overlap = chunk_size / 20; + let (model, chunk_size, chunk_overlap) = Self::config(config)?; let data = RagData::new(&model.id(), chunk_size, chunk_overlap); let mut rag = Self::create(config, name, save_path, data)?; let mut paths = doc_paths.to_vec(); @@ -92,7 +87,7 @@ impl Rag { pub fn create(config: &GlobalConfig, name: &str, path: &Path, data: RagData) -> Result { let hnsw = data.build_hnsw(); - let model = retrieve_embedding_model(&config.read(), &data.model)?; + let model = Model::retrieve_embedding(&config.read(), &data.model)?; let client = init_client(config, Some(model.clone()))?; let rag = Rag { client, @@ -105,6 +100,68 @@ impl Rag { Ok(rag) } + pub fn config(config: &GlobalConfig) -> Result<(Model, usize, usize)> { + let (embedding_model, chunk_size, chunk_overlap) = { + let config = config.read(); + ( + config.rag_embedding_model.clone(), + config.rag_chunk_size, + config.rag_chunk_overlap, + ) + }; + let model_id = match embedding_model { + Some(value) => { + println!("Select embedding model: {value}"); + value + } + None => { + let models = list_embedding_models(&config.read()); + if models.is_empty() { + bail!("No available embedding model"); + } + if *IS_STDOUT_TERMINAL { + select_embedding_model(&models)? + } else { + let value = models[0].id(); + println!("Select embedding model: {value}"); + value + } + } + }; + let model = Model::retrieve_embedding(&config.read(), &model_id)?; + let chunk_size = match chunk_size { + Some(value) => { + println!("Set chunk size: {value}"); + value + } + None => { + if *IS_STDOUT_TERMINAL { + set_chunk_size(&model)? + } else { + let value = model.default_chunk_size(); + println!("Set chunk size: {value}"); + value + } + } + }; + let chunk_overlap = match chunk_overlap { + Some(value) => { + println!("Set chunk overlay: {value}"); + value + } + None => { + let value = chunk_size / 20; + if *IS_STDOUT_TERMINAL { + set_chunk_overlay(value)? + } else { + println!("Set chunk overlay: {value}"); + value + } + } + }; + Ok((model, chunk_size, chunk_overlap)) + } + pub fn save(&self, path: &Path) -> Result<()> { ensure_parent_exists(path)?; let mut file = std::fs::File::create(path)?; @@ -392,27 +449,10 @@ fn document_text(file_path: &str, document: &RagDocument) -> String { ) } -fn retrieve_embedding_model(config: &Config, model_id: &str) -> Result { - let model = Model::find(&list_embedding_models(config), model_id) - .ok_or_else(|| anyhow!("No embedding model '{model_id}'"))?; - Ok(model) -} - -fn select_embedding_model(config: &GlobalConfig) -> Result { - let config = config.read(); - let model = match config.embedding_model.clone() { - Some(model_id) => retrieve_embedding_model(&config, &model_id)?, - None => { - let models = list_embedding_models(&config); - if models.is_empty() { - bail!("No embedding model"); - } - let model_ids: Vec<_> = models.iter().map(|v| v.id()).collect(); - let model_id = Select::new("Select embedding model:", model_ids).prompt()?; - retrieve_embedding_model(&config, &model_id)? - } - }; - Ok(model) +fn select_embedding_model(models: &[&Model]) -> Result { + let model_ids: Vec<_> = models.iter().map(|v| v.id()).collect(); + let model_id = Select::new("Select embedding model:", model_ids).prompt()?; + Ok(model_id) } fn set_chunk_size(model: &Model) -> Result { @@ -437,6 +477,20 @@ fn set_chunk_size(model: &Model) -> Result { value.parse().map_err(|_| anyhow!("Invalid chunk_size")) } +fn set_chunk_overlay(default_value: usize) -> Result { + let value = Text::new("Set chunk overlay:") + .with_default(&default_value.to_string()) + .with_validator(move |text: &str| { + let out = match text.parse::() { + Ok(_) => Validation::Valid, + Err(_) => Validation::Invalid("Must be a integer".into()), + }; + Ok(out) + }) + .prompt()?; + value.parse().map_err(|_| anyhow!("Invalid chunk_overlay")) +} + fn add_doc_paths() -> Result> { let text = Text::new("Add document paths:") .with_validator(required!("This field is required")) -- cgit v1.2.3