summaryrefslogtreecommitdiffstats
path: root/src/rag
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-20 11:26:45 +0800
committerGitHub <noreply@github.com>2024-06-20 11:26:45 +0800
commit2eab71a641827e503b14952373aec82661192ba2 (patch)
treeb2f1fd1d8b988391d85a8058dcaaadd6ef379533 /src/rag
parent3b3d39cef0211b607d51360d0739022f953ba1a7 (diff)
downloadaichat-2eab71a641827e503b14952373aec82661192ba2.tar.gz
feat: rag hybrid search (#618)
Diffstat (limited to 'src/rag')
-rw-r--r--src/rag/bm25.rs172
-rw-r--r--src/rag/mod.rs101
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()
+}