diff options
| author | sigoden <sigoden@gmail.com> | 2024-06-20 11:26:45 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-06-20 11:26:45 +0800 |
| commit | 2eab71a641827e503b14952373aec82661192ba2 (patch) | |
| tree | b2f1fd1d8b988391d85a8058dcaaadd6ef379533 /src/rag | |
| parent | 3b3d39cef0211b607d51360d0739022f953ba1a7 (diff) | |
| download | aichat-2eab71a641827e503b14952373aec82661192ba2.tar.gz | |
feat: rag hybrid search (#618)
Diffstat (limited to 'src/rag')
| -rw-r--r-- | src/rag/bm25.rs | 172 | ||||
| -rw-r--r-- | src/rag/mod.rs | 101 |
2 files changed, 260 insertions, 13 deletions
diff --git a/src/rag/bm25.rs b/src/rag/bm25.rs new file mode 100644 index 0000000..3dee357 --- /dev/null +++ b/src/rag/bm25.rs @@ -0,0 +1,172 @@ +use rayon::prelude::*; +use std::collections::HashMap; +use std::f64; + +#[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.split(' ').map(|v| v.to_string()).collect() +} + +#[cfg(test)] +mod tests { + use super::*; + + #[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 4ce280d..7939d98 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -1,3 +1,4 @@ +use self::bm25::*; use self::loader::*; use self::splitter::*; @@ -5,6 +6,7 @@ use crate::client::*; use crate::config::*; use crate::utils::*; +mod bm25; mod loader; mod splitter; @@ -16,8 +18,7 @@ use inquire::{required, validator::Validation, Select, Text}; use path_absolutize::Absolutize; use serde::{Deserialize, Serialize}; use serde_json::json; -use std::fmt::Debug; -use std::{io::BufReader, path::Path}; +use std::{collections::HashMap, fmt::Debug, io::BufReader, path::Path}; use tokio::sync::mpsc; pub struct Rag { @@ -26,6 +27,7 @@ pub struct Rag { path: String, model: Model, hnsw: Hnsw<'static, f32, DistCosine>, + bm25: BM25<VectorID>, data: RagData, } @@ -85,6 +87,7 @@ impl Rag { pub fn create(config: &GlobalConfig, name: &str, path: &Path, data: RagData) -> Result<Self> { let hnsw = data.build_hnsw(); + let bm25 = data.build_bm25(); let model = Model::retrieve_embedding(&config.read(), &data.model)?; let client = init_client(config, Some(model.clone()))?; let rag = Rag { @@ -94,6 +97,7 @@ impl Rag { data, model, hnsw, + bm25, }; Ok(rag) } @@ -194,12 +198,13 @@ impl Rag { &self, text: &str, top_k: usize, - minimum_score: f32, + min_score_vector: f32, + min_score_text: 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, minimum_score) => { + ret = self.hybird_search(text, top_k, min_score_vector, min_score_text) => { ret } _ = watch_abort_signal(abort_signal) => { @@ -289,18 +294,44 @@ impl Rag { Ok(()) } - async fn search_impl( + async fn hybird_search( &self, - text: &str, + query: &str, top_k: usize, - minimum_score: f32, + min_score_vector: f32, + min_score_text: f32, ) -> Result<Vec<String>> { + let (vector_search_result, text_search_result) = tokio::join!( + self.vector_search(query, top_k, min_score_vector), + self.text_search(query, top_k, min_score_text) + ); + let vector_search_ids = vector_search_result?; + let text_search_ids = text_search_result?; + let ids = reciprocal_rank_fusion(vector_search_ids, text_search_ids, 1.0, 1.0, top_k); + let output: Vec<_> = ids + .into_iter() + .filter_map(|id| { + let (file_index, document_index) = split_vector_id(id); + let file = self.data.files.get(file_index)?; + let document = file.documents.get(document_index)?; + Some(document.page_content.clone()) + }) + .collect(); + Ok(output) + } + + async fn vector_search( + &self, + query: &str, + top_k: usize, + min_score: f32, + ) -> Result<Vec<VectorID>> { let splitter = RecursiveCharacterTextSplitter::new( self.data.chunk_size, self.data.chunk_overlap, &DEFAULT_SEPARATES, ); - let texts = splitter.split_text(text); + let texts = splitter.split_text(query); let embeddings_data = EmbeddingsData::new(texts, true); let embeddings = self.create_embeddings(embeddings_data, None).await?; let output = self @@ -310,13 +341,10 @@ impl Rag { .flat_map(|list| { list.into_iter() .filter_map(|v| { - if v.distance < minimum_score { + if v.distance < min_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.page_content.clone()) + Some(v.d_id) }) .collect::<Vec<_>>() }) @@ -324,6 +352,16 @@ impl Rag { Ok(output) } + async fn text_search( + &self, + query: &str, + top_k: usize, + min_score: f32, + ) -> Result<Vec<VectorID>> { + let output = self.bm25.search(query, top_k, Some(min_score as f64)); + Ok(output) + } + async fn create_embeddings( &self, data: EmbeddingsData, @@ -393,6 +431,17 @@ impl RagData { hnsw.parallel_insert(&list); hnsw } + + pub fn build_bm25(&self) -> BM25<VectorID> { + let mut corpus = vec![]; + for (file_index, file) in self.files.iter().enumerate() { + for (document_index, document) in file.documents.iter().enumerate() { + let id = combine_vector_id(file_index, document_index); + corpus.push((id, document.page_content.clone())); + } + } + BM25::new(corpus, BM25Options::default()) + } } #[derive(Debug, Clone, Serialize, Deserialize)] @@ -502,3 +551,29 @@ fn progress(spinner_message_tx: &Option<mpsc::UnboundedSender<String>>, message: let _ = tx.send(message); } } + +fn reciprocal_rank_fusion( + vector_search_ids: Vec<VectorID>, + text_search_ids: Vec<VectorID>, + vector_search_weight: f32, + text_search_weight: f32, + top_k: usize, +) -> Vec<VectorID> { + let rrf_k = top_k * 2; + let mut map: HashMap<VectorID, f32> = HashMap::new(); + for (index, &item) in vector_search_ids.iter().enumerate() { + *map.entry(item).or_default() += + (1.0 / ((rrf_k + index + 1) as f32)) * vector_search_weight; + } + for (index, &item) in text_search_ids.iter().enumerate() { + *map.entry(item).or_default() += (1.0 / ((rrf_k + index + 1) as f32)) * text_search_weight; + } + let mut sorted_items: Vec<(VectorID, f32)> = map.into_iter().collect(); + sorted_items.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap()); + + sorted_items + .into_iter() + .take(top_k) + .map(|(v, _)| v) + .collect() +} |
