summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-19 12:15:54 +0800
committerGitHub <noreply@github.com>2024-06-19 12:15:54 +0800
commit3b3d39cef0211b607d51360d0739022f953ba1a7 (patch)
tree65fefe139b939ef6f27ac172e8bde89305f51ecc
parent1fb06ecdc4daeed618ec971170e82931464a8399 (diff)
downloadaichat-3b3d39cef0211b607d51360d0739022f953ba1a7.tar.gz
refactor: rag add rag_minimum_score config (#617)
-rw-r--r--README.md1
-rw-r--r--config.example.yaml4
-rw-r--r--models.yaml9
-rw-r--r--src/client/openai.rs3
-rw-r--r--src/config/input.rs9
-rw-r--r--src/config/mod.rs29
-rw-r--r--src/rag/mod.rs27
7 files changed, 52 insertions, 30 deletions
diff --git a/README.md b/README.md
index f7f6e23..0efc4b9 100644
--- a/README.md
+++ b/README.md
@@ -362,6 +362,7 @@ The author mainly focused on writing and programming growing up ...
.set temperature 1.2
.set top_p 0.8
.set rag_top_k 4
+.set rag_minimum_score 0
.set function_calling true
.set compress_threshold 1000
.set dry_run true
diff --git a/config.example.yaml b/config.example.yaml
index 6e4416c..9813ced 100644
--- a/config.example.yaml
+++ b/config.example.yaml
@@ -39,8 +39,10 @@ rag_embedding_model: null
rag_chunk_size: null
# Specifies the chunk overlap
rag_chunk_overlap: null
-# Determines how many relevant documents are retrieved
+# Specifies the number of documents to retrieve
rag_top_k: 4
+# Specifies the minimum relevance score for retrieved documents
+rag_minimum_score: 0
# Defines the query structure using variables like __CONTEXT__ and __INPUT__ to tailor searches to specific needs
rag_template: |
diff --git a/models.yaml b/models.yaml
index 7562f91..1724036 100644
--- a/models.yaml
+++ b/models.yaml
@@ -62,6 +62,11 @@
max_input_tokens: 8191
default_chunk_size: 4000
max_concurrent_chunks: 100
+ - name: text-embedding-ada-002
+ mode: embedding
+ max_input_tokens: 8191
+ default_chunk_size: 4000
+ max_concurrent_chunks: 100
- platform: gemini
# docs:
@@ -609,6 +614,10 @@
input_price: 7
output_price: 7
supports_vision: true
+ - name: embedding-2
+ mode: embedding
+ max_input_tokens: 2048
+ default_chunk_size: 2000
- platform: lingyiwanwu
# docs:
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()?;