summaryrefslogtreecommitdiffstats
path: root/src/rag
diff options
context:
space:
mode:
Diffstat (limited to 'src/rag')
-rw-r--r--src/rag/bm25.rs195
-rw-r--r--src/rag/mod.rs126
-rw-r--r--src/rag/serde_vectors.rs6
3 files changed, 74 insertions, 253 deletions
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}'")))?;