diff options
| author | sigoden <sigoden@gmail.com> | 2024-10-17 19:36:00 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-10-17 19:36:00 +0800 |
| commit | db3a7b6ec2dee82df74fa055ab83b5a140cf49f0 (patch) | |
| tree | 919f258da08eee675e2cefb30c31b3c2f05dd4b2 /src | |
| parent | af1e57f7c6e3259be48783175289b158cd763a10 (diff) | |
| download | aichat-db3a7b6ec2dee82df74fa055ab83b5a140cf49f0.tar.gz | |
refactor: improve RAG search (#931)
Diffstat (limited to 'src')
| -rw-r--r-- | src/config/mod.rs | 23 | ||||
| -rw-r--r-- | src/rag/bm25.rs | 195 | ||||
| -rw-r--r-- | src/rag/mod.rs | 126 | ||||
| -rw-r--r-- | src/rag/serde_vectors.rs | 6 |
4 files changed, 88 insertions, 262 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs index 34da912..0f3f89e 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -65,19 +65,24 @@ 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#"Use the following context as your learned knowledge, inside <context></context> XML tags. +const RAG_TEMPLATE: &str = r#"Answer the query based on the context while respecting the rules. (user query, some textual context and rules, all inside xml tags) + <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. - -Given the context information, answer the query. -Query: __INPUT__"#; +<rules> +- If you don't know, just say so. +- If you are not sure, ask for clarification. +- Answer in the same language as the user query. +- If the context appears unreadable or of poor quality, tell the user then answer as best as you can. +- If the answer is not in the context but you think you know the answer, explain that to the user then answer with your own knowledge. +- Answer directly and without using xml tags. +</rules> + +<user_query> +__INPUT__ +</user_query>"#; const LEFT_PROMPT: &str = "{color.green}{?session {?agent {agent}>}{session}{?role /}}{!session {?agent {agent}>}}{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}"; diff --git a/src/rag/bm25.rs b/src/rag/bm25.rs deleted file mode 100644 index c91dcf5..0000000 --- a/src/rag/bm25.rs +++ /dev/null @@ -1,195 +0,0 @@ -use rayon::prelude::*; -use std::collections::HashMap; -use std::f64; -use unicode_segmentation::UnicodeSegmentation; - -#[derive(Debug, Clone)] -pub struct BM25Options { - k1: f64, - b: f64, - epsilon: f64, -} - -impl Default for BM25Options { - fn default() -> Self { - Self { - k1: 1.5, - b: 0.75, - epsilon: 0.25, - } - } -} - -#[derive(Debug, Clone)] -pub struct BM25<T> { - options: BM25Options, - corpus_size: usize, - avgdl: f64, - doc_freqs: Vec<HashMap<String, u32>>, - doc_ids: Vec<T>, - idf: HashMap<String, f64>, - doc_len: Vec<usize>, -} - -impl<T: Clone> BM25<T> { - pub fn new(corpus: Vec<(T, String)>, options: BM25Options) -> Self { - let mut doc_ids = vec![]; - let mut docs = vec![]; - for (id, value) in corpus { - doc_ids.push(id); - docs.push(value); - } - let tokenized_docs = docs.into_par_iter().map(|text| tokenize(&text)).collect(); - - let mut bm25 = BM25 { - options, - corpus_size: 0, - avgdl: 0.0, - doc_freqs: Vec::new(), - doc_ids, - idf: HashMap::new(), - doc_len: Vec::new(), - }; - - let map = bm25.initialize(tokenized_docs); - bm25.calc_idf(map); - - bm25 - } - - pub fn search(&self, query: &str, top_k: usize, min_score: Option<f64>) -> Vec<T> { - let scores = self.get_scores(query); - let mut indexed_scores: Vec<(T, f64)> = scores - .into_iter() - .enumerate() - .filter_map(|(i, v)| match min_score { - Some(minimum_score) => { - if v < minimum_score { - None - } else { - Some((self.doc_ids[i].clone(), v)) - } - } - None => Some((self.doc_ids[i].clone(), v)), - }) - .collect(); - indexed_scores.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap()); - indexed_scores - .into_iter() - .take(top_k) - .map(|(id, _)| id) - .collect() - } - - pub fn get_scores(&self, query: &str) -> Vec<f64> { - let mut score = vec![0.0; self.corpus_size]; - - for q in tokenize(query) { - if let Some(idf) = self.idf.get(&q) { - for (i, doc) in self.doc_freqs.iter().enumerate() { - let q_freq = doc.get(&q).unwrap_or(&0); - score[i] += *idf - * (*q_freq as f64 * (self.options.k1 + 1.0) - / (*q_freq as f64 - + self.options.k1 - * (1.0 - self.options.b - + self.options.b * self.doc_len[i] as f64 / self.avgdl))); - } - } - } - - score - } - - fn initialize(&mut self, corpus: Vec<Vec<String>>) -> HashMap<String, usize> { - let mut map = HashMap::new(); - let mut num_doc = 0; - - for document in corpus { - self.doc_len.push(document.len()); - num_doc += document.len(); - - let mut frequencies = HashMap::new(); - for word in document { - *frequencies.entry(word).or_insert(0) += 1; - } - self.doc_freqs.push(frequencies); - - for word in self.doc_freqs[self.doc_freqs.len() - 1].keys() { - *map.entry(word.clone()).or_insert(0) += 1; - } - - self.corpus_size += 1; - } - - self.avgdl = num_doc as f64 / self.corpus_size as f64; - map - } - - fn calc_idf(&mut self, map: HashMap<String, usize>) { - let mut idf_sum = 0.0; - let mut negative_idfs = Vec::new(); - - for (word, freq) in map { - let idf = (self.corpus_size as f64 - freq as f64 + 0.5).ln() - (freq as f64 + 0.5).ln(); - self.idf.insert(word.clone(), idf); - idf_sum += idf; - if idf < 0.0 { - negative_idfs.push(word); - } - } - - let average_idf = idf_sum / self.idf.len() as f64; - - for word in negative_idfs { - self.idf.insert(word, self.options.epsilon * average_idf); - } - } -} - -fn tokenize(text: &str) -> Vec<String> { - text.unicode_words() - .filter_map(|v| { - if [ - "a", "an", "and", "are", "as", "at", "be", "but", "by", "for", "if", "in", "into", - "is", "it", "no", "not", "of", "on", "or", "such", "that", "the", "their", "then", - "there", "these", "they", "this", "to", "was", "will", "with", - ] - .contains(&v) - { - None - } else { - Some(v.to_string()) - } - }) - .collect() -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_tokenize() { - assert_eq!( - tokenize("a quick fox jumps over the lazy dog"), - vec!["quick", "fox", "jumps", "over", "lazy", "dog"] - ); - } - - #[test] - fn test_bm25() { - let corpus = vec![ - (0, "Hello there good man!".into()), - (1, "It is quite windy in London".into()), - (2, "How is the weather today?".into()), - ]; - let bm25 = BM25::new(corpus, BM25Options::default()); - - let scores = bm25.get_scores("windy London"); - assert_eq!(scores, [0.0, 0.9372947225064051, 0.0]); - - let top_n = bm25.search("windy London", 3, None); - assert_eq!(top_n, vec![1, 0, 2]) - } -} diff --git a/src/rag/mod.rs b/src/rag/mod.rs index e84ece3..4cee391 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -1,4 +1,3 @@ -use self::bm25::*; use self::loader::*; use self::splitter::*; @@ -6,12 +5,12 @@ use crate::client::*; use crate::config::*; use crate::utils::*; -mod bm25; mod loader; mod serde_vectors; mod splitter; use anyhow::{anyhow, bail, Context, Result}; +use bm25::{LanguageMode, SearchEngine, SearchEngineBuilder}; use hnsw_rs::prelude::*; use indexmap::{IndexMap, IndexSet}; use inquire::{required, validator::Validation, Confirm, Select, Text}; @@ -19,7 +18,7 @@ use parking_lot::RwLock; use path_absolutize::Absolutize; use serde::{Deserialize, Serialize}; use serde_json::json; -use std::{collections::HashMap, env, fmt::Debug, fs, path::Path, time::Duration}; +use std::{collections::HashMap, env, fmt::Debug, fs, hash::Hash, path::Path, time::Duration}; use tokio::time::sleep; pub struct Rag { @@ -28,7 +27,7 @@ pub struct Rag { path: String, embedding_model: Model, hnsw: Hnsw<'static, f32, DistCosine>, - bm25: BM25<DocumentId>, + bm25: SearchEngine<DocumentId>, data: RagData, last_sources: RwLock<Option<String>>, } @@ -52,7 +51,7 @@ impl Clone for Rag { path: self.path.clone(), embedding_model: self.embedding_model.clone(), hnsw: self.data.build_hnsw(), - bm25: self.bm25.clone(), + bm25: self.data.build_bm25(), data: self.data.clone(), last_sources: RwLock::new(None), } @@ -230,7 +229,7 @@ impl Rag { let sources: IndexSet<_> = ids .iter() .filter_map(|id| { - let (file_index, _) = split_document_id(*id); + let (file_index, _) = id.split(); let file = self.data.files.get(&file_index)?; Some(file.path.clone()) }) @@ -394,14 +393,7 @@ impl Rag { &separator, ); - let metadata = metadata - .iter() - .map(|(k, v)| format!("{k}: {v}\n")) - .collect::<Vec<String>>() - .join(""); - let split_options = SplitterChunkHeaderOptions::default().with_chunk_header(&format!( - "<document_metadata>\npath: {path}\n{metadata}</document_metadata>\n\n" - )); + let split_options = SplitterChunkHeaderOptions::default(); let document = RagDocument::new(contents); let split_documents = splitter.split_documents(&[document], &split_options); rag_files.push(RagFile { @@ -420,7 +412,7 @@ impl Rag { let mut texts = vec![]; for file in rag_files.into_iter() { for (document_index, document) in file.documents.iter().enumerate() { - document_ids.push(combine_document_id(next_file_id, document_index)); + document_ids.push(DocumentId::new(next_file_id, document_index)); texts.push(document.page_content.clone()) } files.push((next_file_id, file)); @@ -456,17 +448,21 @@ impl Rag { min_score_keyword_search: f32, rerank_model: Option<&str>, ) -> Result<Vec<(DocumentId, String)>> { - let (vector_search_result, text_search_result) = tokio::join!( + let (vector_search_results, keyword_search_results) = tokio::join!( self.vector_search(query, top_k, min_score_vector_search), self.keyword_search(query, top_k, min_score_keyword_search) ); - let vector_search_ids = vector_search_result?; - let keyword_search_ids = text_search_result?; - debug!( - "vector_search_ids: {:?}, keyword_search_ids: {:?}", - pretty_document_ids(&vector_search_ids), - pretty_document_ids(&keyword_search_ids) - ); + + let vector_search_results = vector_search_results?; + debug!("vector_search_results: {vector_search_results:?}",); + let vector_search_ids: Vec<DocumentId> = + vector_search_results.into_iter().map(|(v, _)| v).collect(); + + let keyword_search_results = keyword_search_results?; + debug!("keyword_search_results: {keyword_search_results:?}",); + let keyword_search_ids: Vec<DocumentId> = + keyword_search_results.into_iter().map(|(v, _)| v).collect(); + let ids = match rerank_model { Some(model_id) => { let model = Model::retrieve_reranker(&self.config.read(), model_id)?; @@ -490,7 +486,7 @@ impl Rag { .take(top_k) .filter_map(|item| documents_ids.get(item.index).cloned()) .collect(); - debug!("rerank_ids: {:?}", pretty_document_ids(&ids)); + debug!("rerank_ids: {ids:?}"); ids } None => { @@ -499,7 +495,7 @@ impl Rag { vec![1.0, 1.0], top_k, ); - debug!("rrf_ids: {:?}", pretty_document_ids(&ids)); + debug!("rrf_ids: {ids:?}"); ids } }; @@ -518,7 +514,7 @@ impl Rag { query: &str, top_k: usize, min_score: f32, - ) -> Result<Vec<DocumentId>> { + ) -> Result<Vec<(DocumentId, f32)>> { let splitter = RecursiveCharacterTextSplitter::new( self.data.chunk_size, self.data.chunk_overlap, @@ -534,10 +530,12 @@ impl Rag { .flat_map(|list| { list.into_iter() .filter_map(|v| { - if v.distance < min_score { - return None; + let score = 1.0 - v.distance; + if score > min_score { + Some((DocumentId(v.d_id), score)) + } else { + None } - Some(v.d_id) }) .collect::<Vec<_>>() }) @@ -550,8 +548,19 @@ impl Rag { query: &str, top_k: usize, min_score: f32, - ) -> Result<Vec<DocumentId>> { - let output = self.bm25.search(query, top_k, Some(min_score as f64)); + ) -> Result<Vec<(DocumentId, f32)>> { + let results = self.bm25.search(query, top_k); + let output: Vec<(DocumentId, f32)> = results + .into_iter() + .filter_map(|v| { + let score = v.score; + if score > min_score { + Some((v.document.id, score)) + } else { + None + } + }) + .collect(); Ok(output) } @@ -670,7 +679,7 @@ impl RagData { } pub fn get(&self, id: DocumentId) -> Option<&RagDocument> { - let (file_index, document_index) = split_document_id(id); + let (file_index, document_index) = id.split(); let file = self.files.get(&file_index)?; let document = file.documents.get(document_index)?; Some(document) @@ -680,7 +689,7 @@ impl RagData { for file_id in file_ids { if let Some(file) = self.files.swap_remove(&file_id) { for (document_index, _) in file.documents.iter().enumerate() { - let document_id = combine_document_id(file_id, document_index); + let document_id = DocumentId::new(file_id, document_index); self.vectors.swap_remove(&document_id); } } @@ -702,20 +711,23 @@ impl RagData { pub fn build_hnsw(&self) -> Hnsw<'static, f32, DistCosine> { let hnsw = Hnsw::new(32, self.vectors.len(), 16, 200, DistCosine {}); - let list: Vec<_> = self.vectors.iter().map(|(k, v)| (v, *k)).collect(); + let list: Vec<_> = self.vectors.iter().map(|(k, v)| (v, k.0)).collect(); hnsw.parallel_insert(&list); hnsw } - pub fn build_bm25(&self) -> BM25<DocumentId> { - let mut corpus = vec![]; + pub fn build_bm25(&self) -> SearchEngine<DocumentId> { + let mut documents = vec![]; for (file_index, file) in self.files.iter() { for (document_index, document) in file.documents.iter().enumerate() { - let id = combine_document_id(*file_index, document_index); - corpus.push((id, document.page_content.clone())); + let id = DocumentId::new(*file_index, document_index); + documents.push(bm25::Document::new(id, &document.page_content)) } } - BM25::new(corpus, BM25Options::default()) + SearchEngineBuilder::<DocumentId>::with_documents(LanguageMode::Detect, documents) + .k1(1.5) + .b(0.75) + .build() } } @@ -753,26 +765,30 @@ impl Default for RagDocument { pub type RagMetadata = IndexMap<String, String>; pub type FileId = usize; -pub type DocumentId = usize; -pub fn combine_document_id(file_index: usize, document_index: usize) -> DocumentId { - file_index << (usize::BITS / 2) | document_index -} +#[derive(Clone, Copy, Hash, Eq, PartialEq, Ord, PartialOrd)] +pub struct DocumentId(usize); -pub fn split_document_id(value: DocumentId) -> (usize, usize) { - let low_mask = (1 << (usize::BITS / 2)) - 1; - let low = value & low_mask; - let high = value >> (usize::BITS / 2); - (high, low) +impl Debug for DocumentId { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + let (file_index, document_index) = self.split(); + f.write_fmt(format_args!("{file_index}-{document_index}")) + } } -fn pretty_document_ids(ids: &[DocumentId]) -> Vec<String> { - ids.iter() - .map(|v| { - let (h, l) = split_document_id(*v); - format!("{h}-{l}") - }) - .collect() +impl DocumentId { + pub fn new(file_index: usize, document_index: usize) -> Self { + let value = file_index << (usize::BITS / 2) | document_index; + Self(value) + } + + pub fn split(self) -> (usize, usize) { + let value = self.0; + let low_mask = (1 << (usize::BITS / 2)) - 1; + let low = value & low_mask; + let high = value >> (usize::BITS / 2); + (high, low) + } } fn select_embedding_model(models: &[&Model]) -> Result<String> { diff --git a/src/rag/serde_vectors.rs b/src/rag/serde_vectors.rs index 894c22c..bd821fa 100644 --- a/src/rag/serde_vectors.rs +++ b/src/rag/serde_vectors.rs @@ -12,8 +12,8 @@ where { let encoded_map: IndexMap<String, String> = vectors .iter() - .map(|(key, vec)| { - let (h, l) = split_document_id(*key); + .map(|(id, vec)| { + let (h, l) = id.split(); let byte_slice = unsafe { std::slice::from_raw_parts( vec.as_ptr() as *const u8, @@ -41,7 +41,7 @@ where .and_then(|(h, l)| { let h = h.parse::<usize>().ok()?; let l = l.parse::<usize>().ok()?; - Some(combine_document_id(h, l)) + Some(DocumentId::new(h, l)) }) .ok_or_else(|| de::Error::custom(format!("Invalid key '{key}'")))?; |
