diff options
| author | sigoden <sigoden@gmail.com> | 2024-06-19 12:15:54 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-06-19 12:15:54 +0800 |
| commit | 3b3d39cef0211b607d51360d0739022f953ba1a7 (patch) | |
| tree | 65fefe139b939ef6f27ac172e8bde89305f51ecc /src | |
| parent | 1fb06ecdc4daeed618ec971170e82931464a8399 (diff) | |
| download | aichat-3b3d39cef0211b607d51360d0739022f953ba1a7.tar.gz | |
refactor: rag add rag_minimum_score config (#617)
Diffstat (limited to 'src')
| -rw-r--r-- | src/client/openai.rs | 3 | ||||
| -rw-r--r-- | src/config/input.rs | 9 | ||||
| -rw-r--r-- | src/config/mod.rs | 29 | ||||
| -rw-r--r-- | src/rag/mod.rs | 27 |
4 files changed, 39 insertions, 29 deletions
diff --git a/src/client/openai.rs b/src/client/openai.rs index 9d616ed..0c51b33 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -241,8 +241,7 @@ pub fn openai_build_chat_completions_body(data: ChatCompletionsData, model: &Mod pub fn openai_build_embeddings_body(data: EmbeddingsData, model: &Model) -> Value { json!({ "input": data.texts, - "model": model.name(), - "encoding_format": "float", + "model": model.name() }) } diff --git a/src/config/input.rs b/src/config/input.rs index 4ddc06f..935d517 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -169,8 +169,13 @@ impl Input { if !self.text.is_empty() { let rag = self.config.read().rag.clone(); if let Some(rag) = rag { - let top_k = self.config.read().rag_top_k; - let embeddings = rag.search(&self.text, top_k, abort_signal).await?; + let (top_k, minimum_score) = { + let config = self.config.read(); + (config.rag_top_k, config.rag_minimum_score) + }; + let embeddings = rag + .search(&self.text, top_k, minimum_score, abort_signal) + .await?; let text = self.config.read().rag_template(&embeddings, &self.text); self.patched_text = Some(text); self.rag_name = Some(rag.name().to_string()); diff --git a/src/config/mod.rs b/src/config/mod.rs index 41e449a..b2e4517 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -62,18 +62,18 @@ const SUMMARIZE_PROMPT: &str = const SUMMARY_PROMPT: &str = "This is a summary of the chat history as a recap: "; const RAG_TEMPLATE: &str = r#"Use the following context as your learned knowledge, inside <context></context> XML tags. - <context> - __CONTEXT__ - </context> +<context> +__CONTEXT__ +</context> - 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. +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__"#; +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}"; @@ -105,6 +105,7 @@ pub struct Config { pub rag_chunk_size: Option<usize>, pub rag_chunk_overlap: Option<usize>, pub rag_top_k: usize, + pub rag_minimum_score: f32, pub rag_template: Option<String>, pub compress_threshold: usize, pub summarize_prompt: Option<String>, @@ -157,6 +158,7 @@ impl Default for Config { rag_chunk_size: None, rag_chunk_overlap: None, rag_top_k: 4, + rag_minimum_score: 0.0, rag_template: None, compress_threshold: 4000, summarize_prompt: None, @@ -439,6 +441,7 @@ impl Config { ("temperature", format_option_value(&role.temperature())), ("top_p", format_option_value(&role.top_p())), ("rag_top_k", self.rag_top_k.to_string()), + ("rag_minimum_score", self.rag_minimum_score.to_string()), ("function_calling", self.function_calling.to_string()), ("compress_threshold", self.compress_threshold.to_string()), ("dry_run", self.dry_run.to_string()), @@ -491,6 +494,11 @@ impl Config { self.rag_top_k = value; } } + "rag_minimum_score" => { + if let Some(value) = parse_value(value)? { + self.rag_minimum_score = value; + } + } "function_calling" => { let value = value.parse().with_context(|| "Invalid value")?; self.function_calling = value; @@ -1049,6 +1057,7 @@ impl Config { "temperature", "top_p", "rag_top_k", + "rag_minimum_score", "function_calling", "compress_threshold", "save", diff --git a/src/rag/mod.rs b/src/rag/mod.rs index b7ac6cd..4ce280d 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -20,8 +20,6 @@ use std::fmt::Debug; use std::{io::BufReader, path::Path}; use tokio::sync::mpsc; -pub const SIMILARITY_THRESHOLD: f32 = 0.25; - pub struct Rag { client: Box<dyn Client>, name: String, @@ -176,6 +174,7 @@ impl Rag { "path": self.path, "model": self.model.id(), "chunk_size": self.data.chunk_size, + "chunk_overlap": self.data.chunk_overlap, "files": files, }); let output = serde_yaml::to_string(&data) @@ -195,11 +194,12 @@ impl Rag { &self, text: &str, top_k: usize, + minimum_score: f32, abort_signal: AbortSignal, ) -> Result<String> { let (stop_spinner_tx, _) = run_spinner("Searching").await; let ret = tokio::select! { - ret = self.search_impl(text, top_k) => { + ret = self.search_impl(text, top_k, minimum_score) => { ret } _ = watch_abort_signal(abort_signal) => { @@ -273,7 +273,7 @@ impl Rag { for (file_index, file) in rag_files.iter().enumerate() { for (document_index, document) in file.documents.iter().enumerate() { vector_ids.push(combine_vector_id(file_index, document_index)); - texts.push(document_text(&file.path, document)) + texts.push(document.page_content.clone()) } } @@ -289,7 +289,12 @@ impl Rag { Ok(()) } - async fn search_impl(&self, text: &str, top_k: usize) -> Result<Vec<String>> { + async fn search_impl( + &self, + text: &str, + top_k: usize, + minimum_score: f32, + ) -> Result<Vec<String>> { let splitter = RecursiveCharacterTextSplitter::new( self.data.chunk_size, self.data.chunk_overlap, @@ -305,13 +310,13 @@ impl Rag { .flat_map(|list| { list.into_iter() .filter_map(|v| { - if v.distance < SIMILARITY_THRESHOLD { + if v.distance < minimum_score { return None; } let (file_index, document_index) = split_vector_id(v.d_id); let file = self.data.files.get(file_index)?; let document = file.documents.get(document_index)?; - Some(document_text(&file.path, document)) + Some(document.page_content.clone()) }) .collect::<Vec<_>>() }) @@ -441,14 +446,6 @@ pub fn split_vector_id(value: VectorID) -> (usize, usize) { (high, low) } -fn document_text(file_path: &str, document: &RagDocument) -> String { - format!( - "file_path: {}\n\n{}", - shell_words::quote(file_path), - document.page_content - ) -} - fn select_embedding_model(models: &[&Model]) -> Result<String> { let model_ids: Vec<_> = models.iter().map(|v| v.id()).collect(); let model_id = Select::new("Select embedding model:", model_ids).prompt()?; |
