summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-21 07:22:43 +0800
committerGitHub <noreply@github.com>2024-06-21 07:22:43 +0800
commit34d568d5e229729031245aee488481c7f623125c (patch)
treeb54f078561c4ffe9fb7ae445788b0b495b6e0450 /src
parentba832016f3b272aac561ac9b053fcb9df1cb8d13 (diff)
downloadaichat-34d568d5e229729031245aee488481c7f623125c.tar.gz
refactor: improve rag (#622)
Diffstat (limited to 'src')
-rw-r--r--src/config/input.rs6
-rw-r--r--src/config/mod.rs4
-rw-r--r--src/rag/bm25.rs25
-rw-r--r--src/rag/mod.rs81
4 files changed, 71 insertions, 45 deletions
diff --git a/src/config/input.rs b/src/config/input.rs
index 03ec999..b9b1746 100644
--- a/src/config/input.rs
+++ b/src/config/input.rs
@@ -169,12 +169,12 @@ impl Input {
if !self.text.is_empty() {
let rag = self.config.read().rag.clone();
if let Some(rag) = rag {
- let (top_k, min_score_vector_search, min_score_fulltext_search) = {
+ let (top_k, min_score_vector_search, min_score_keyword_search) = {
let config = self.config.read();
(
config.rag_top_k,
config.rag_min_score_vector_search,
- config.rag_min_score_fulltext_search,
+ config.rag_min_score_keyword_search,
)
};
let rerank = match self.config.read().rag_rerank_model.clone() {
@@ -192,7 +192,7 @@ impl Input {
&self.text,
top_k,
min_score_vector_search,
- min_score_fulltext_search,
+ min_score_keyword_search,
rerank,
abort_signal,
)
diff --git a/src/config/mod.rs b/src/config/mod.rs
index e07dca6..903face 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -107,7 +107,7 @@ pub struct Config {
pub rag_chunk_overlap: Option<usize>,
pub rag_top_k: usize,
pub rag_min_score_vector_search: f32,
- pub rag_min_score_fulltext_search: f32,
+ pub rag_min_score_keyword_search: f32,
pub rag_min_score_rerank: f32,
pub rag_template: Option<String>,
pub compress_threshold: usize,
@@ -163,7 +163,7 @@ impl Default for Config {
rag_chunk_overlap: None,
rag_top_k: 4,
rag_min_score_vector_search: 0.0,
- rag_min_score_fulltext_search: 0.0,
+ rag_min_score_keyword_search: 0.0,
rag_min_score_rerank: 0.0,
rag_template: None,
compress_threshold: 4000,
diff --git a/src/rag/bm25.rs b/src/rag/bm25.rs
index 3dee357..c91dcf5 100644
--- a/src/rag/bm25.rs
+++ b/src/rag/bm25.rs
@@ -1,6 +1,7 @@
use rayon::prelude::*;
use std::collections::HashMap;
use std::f64;
+use unicode_segmentation::UnicodeSegmentation;
#[derive(Debug, Clone)]
pub struct BM25Options {
@@ -147,7 +148,21 @@ impl<T: Clone> BM25<T> {
}
fn tokenize(text: &str) -> Vec<String> {
- text.split(' ').map(|v| v.to_string()).collect()
+ 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)]
@@ -155,6 +170,14 @@ 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()),
diff --git a/src/rag/mod.rs b/src/rag/mod.rs
index 16c3d99..628a9e4 100644
--- a/src/rag/mod.rs
+++ b/src/rag/mod.rs
@@ -13,13 +13,12 @@ mod splitter;
use anyhow::bail;
use anyhow::{anyhow, Context, Result};
use hnsw_rs::prelude::*;
-use indexmap::IndexMap;
-use indexmap::IndexSet;
+use indexmap::{IndexMap, IndexSet};
use inquire::{required, validator::Validation, Select, Text};
use path_absolutize::Absolutize;
use serde::{Deserialize, Serialize};
use serde_json::json;
-use std::{collections::HashMap, fmt::Debug, io::BufReader, path::Path};
+use std::{fmt::Debug, io::BufReader, path::Path};
use tokio::sync::mpsc;
pub struct Rag {
@@ -27,7 +26,7 @@ pub struct Rag {
path: String,
embedding_model: Model,
hnsw: Hnsw<'static, f32, DistCosine>,
- bm25: BM25<VectorID>,
+ bm25: BM25<DocumentId>,
data: RagData,
embedding_client: Box<dyn Client>,
}
@@ -202,13 +201,13 @@ impl Rag {
text: &str,
top_k: usize,
min_score_vector_search: f32,
- min_score_fulltext_search: f32,
+ min_score_keyword_search: f32,
rerank: Option<(Box<dyn Client>, f32)>,
abort_signal: AbortSignal,
) -> Result<String> {
let (stop_spinner_tx, _) = run_spinner("Searching").await;
let ret = tokio::select! {
- ret = self.hybird_search(text, top_k, min_score_vector_search, min_score_fulltext_search, rerank) => {
+ ret = self.hybird_search(text, top_k, min_score_vector_search, min_score_keyword_search, rerank) => {
ret
}
_ = watch_abort_signal(abort_signal) => {
@@ -281,7 +280,7 @@ impl Rag {
let mut texts = vec![];
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));
+ vector_ids.push(combine_document_id(file_index, document_index));
texts.push(document.page_content.clone())
}
}
@@ -304,20 +303,22 @@ impl Rag {
query: &str,
top_k: usize,
min_score_vector_search: f32,
- min_score_fulltext_search: f32,
+ min_score_keyword_search: f32,
rerank: Option<(Box<dyn Client>, f32)>,
) -> Result<Vec<String>> {
let (vector_search_result, text_search_result) = tokio::join!(
self.vector_search(query, top_k, min_score_vector_search),
- self.fulltext_search(query, top_k, min_score_fulltext_search)
+ self.keyword_search(query, top_k, min_score_keyword_search)
);
let vector_search_ids = vector_search_result?;
- let fulltext_search_ids = text_search_result?;
- debug!("vector_search_ids: {vector_search_ids:?}, fulltext_search_ids: {fulltext_search_ids:?}");
+ let keyword_search_ids = text_search_result?;
+ debug!(
+ "vector_search_ids: {vector_search_ids:?}, keyword_search_ids: {keyword_search_ids:?}"
+ );
let ids = match rerank {
Some((client, min_score)) => {
let min_score = min_score as f64;
- let ids: IndexSet<VectorID> = [vector_search_ids, fulltext_search_ids]
+ let ids: IndexSet<DocumentId> = [vector_search_ids, keyword_search_ids]
.concat()
.into_iter()
.collect();
@@ -345,8 +346,11 @@ impl Rag {
ids
}
None => {
- let ids =
- reciprocal_rank_fusion(vector_search_ids, fulltext_search_ids, 1.0, 1.0, top_k);
+ let ids = reciprocal_rank_fusion(
+ vec![vector_search_ids, keyword_search_ids],
+ vec![1.0, 1.0],
+ top_k,
+ );
debug!("rrf_ids: {ids:?}");
ids
}
@@ -366,7 +370,7 @@ impl Rag {
query: &str,
top_k: usize,
min_score: f32,
- ) -> Result<Vec<VectorID>> {
+ ) -> Result<Vec<DocumentId>> {
let splitter = RecursiveCharacterTextSplitter::new(
self.data.chunk_size,
self.data.chunk_overlap,
@@ -393,12 +397,12 @@ impl Rag {
Ok(output)
}
- async fn fulltext_search(
+ async fn keyword_search(
&self,
query: &str,
top_k: usize,
min_score: f32,
- ) -> Result<Vec<VectorID>> {
+ ) -> Result<Vec<DocumentId>> {
let output = self.bm25.search(query, top_k, Some(min_score as f64));
Ok(output)
}
@@ -442,7 +446,7 @@ pub struct RagData {
pub chunk_size: usize,
pub chunk_overlap: usize,
pub files: Vec<RagFile>,
- pub vectors: IndexMap<VectorID, Vec<f32>>,
+ pub vectors: IndexMap<DocumentId, Vec<f32>>,
}
impl RagData {
@@ -456,8 +460,8 @@ impl RagData {
}
}
- pub fn get(&self, id: VectorID) -> Option<&RagDocument> {
- let (file_index, document_index) = split_vector_id(id);
+ pub fn get(&self, id: DocumentId) -> Option<&RagDocument> {
+ let (file_index, document_index) = split_document_id(id);
let file = self.files.get(file_index)?;
let document = file.documents.get(document_index)?;
Some(document)
@@ -466,7 +470,7 @@ impl RagData {
pub fn add(
&mut self,
files: Vec<RagFile>,
- vector_ids: Vec<VectorID>,
+ vector_ids: Vec<DocumentId>,
embeddings: EmbeddingsOutput,
) {
self.files.extend(files);
@@ -480,11 +484,11 @@ impl RagData {
hnsw
}
- pub fn build_bm25(&self) -> BM25<VectorID> {
+ pub fn build_bm25(&self) -> BM25<DocumentId> {
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);
+ let id = combine_document_id(file_index, document_index);
corpus.push((id, document.page_content.clone()));
}
}
@@ -530,13 +534,13 @@ impl Default for RagDocument {
pub type RagMetadata = IndexMap<String, String>;
-pub type VectorID = usize;
+pub type DocumentId = usize;
-pub fn combine_vector_id(file_index: usize, document_index: usize) -> VectorID {
+pub fn combine_document_id(file_index: usize, document_index: usize) -> DocumentId {
file_index << (usize::BITS / 2) | document_index
}
-pub fn split_vector_id(value: VectorID) -> (usize, 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);
@@ -601,22 +605,21 @@ fn progress(spinner_message_tx: &Option<mpsc::UnboundedSender<String>>, message:
}
fn reciprocal_rank_fusion(
- vector_search_ids: Vec<VectorID>,
- text_search_ids: Vec<VectorID>,
- vector_search_weight: f32,
- text_search_weight: f32,
+ list_of_document_ids: Vec<Vec<DocumentId>>,
+ list_of_weights: Vec<f32>,
top_k: usize,
-) -> Vec<VectorID> {
+) -> Vec<DocumentId> {
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 map: IndexMap<DocumentId, f32> = IndexMap::new();
+ for (document_ids, weight) in list_of_document_ids
+ .into_iter()
+ .zip(list_of_weights.into_iter())
+ {
+ for (index, &item) in document_ids.iter().enumerate() {
+ *map.entry(item).or_default() += (1.0 / ((rrf_k + index + 1) as f32)) * weight;
+ }
}
- let mut sorted_items: Vec<(VectorID, f32)> = map.into_iter().collect();
+ let mut sorted_items: Vec<(DocumentId, f32)> = map.into_iter().collect();
sorted_items.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap());
sorted_items