diff options
Diffstat (limited to 'src/rag/mod.rs')
| -rw-r--r-- | src/rag/mod.rs | 118 |
1 files changed, 83 insertions, 35 deletions
diff --git a/src/rag/mod.rs b/src/rag/mod.rs index 7939d98..16c3d99 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -14,6 +14,7 @@ use anyhow::bail; use anyhow::{anyhow, Context, Result}; use hnsw_rs::prelude::*; use indexmap::IndexMap; +use indexmap::IndexSet; use inquire::{required, validator::Validation, Select, Text}; use path_absolutize::Absolutize; use serde::{Deserialize, Serialize}; @@ -22,13 +23,13 @@ use std::{collections::HashMap, fmt::Debug, io::BufReader, path::Path}; use tokio::sync::mpsc; pub struct Rag { - client: Box<dyn Client>, name: String, path: String, - model: Model, + embedding_model: Model, hnsw: Hnsw<'static, f32, DistCosine>, bm25: BM25<VectorID>, data: RagData, + embedding_client: Box<dyn Client>, } impl Debug for Rag { @@ -36,7 +37,7 @@ impl Debug for Rag { f.debug_struct("Rag") .field("name", &self.name) .field("path", &self.path) - .field("model", &self.model) + .field("embedding_model", &self.embedding_model) .field("data", &self.data) .finish() } @@ -51,8 +52,8 @@ impl Rag { abort_signal: AbortSignal, ) -> Result<Self> { debug!("init rag: {name}"); - let (model, chunk_size, chunk_overlap) = Self::config(config)?; - let data = RagData::new(&model.id(), chunk_size, chunk_overlap); + let (embedding_model, chunk_size, chunk_overlap) = Self::config(config)?; + let data = RagData::new(embedding_model.id(), chunk_size, chunk_overlap); let mut rag = Self::create(config, name, save_path, data)?; let mut paths = doc_paths.to_vec(); if paths.is_empty() { @@ -88,22 +89,22 @@ 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 embedding_model = Model::retrieve_embedding(&config.read(), &data.embedding_model)?; + let embedding_client = init_client(config, Some(embedding_model.clone()))?; let rag = Rag { - client, name: name.to_string(), path: path.display().to_string(), data, - model, + embedding_model, hnsw, bm25, + embedding_client, }; Ok(rag) } pub fn config(config: &GlobalConfig) -> Result<(Model, usize, usize)> { - let (embedding_model, chunk_size, chunk_overlap) = { + let (embedding_model_id, chunk_size, chunk_overlap) = { let config = config.read(); ( config.rag_embedding_model.clone(), @@ -111,7 +112,7 @@ impl Rag { config.rag_chunk_overlap, ) }; - let model_id = match embedding_model { + let embedding_model_id = match embedding_model_id { Some(value) => { println!("Select embedding model: {value}"); value @@ -130,7 +131,8 @@ impl Rag { } } }; - let model = Model::retrieve_embedding(&config.read(), &model_id)?; + let embedding_model = Model::retrieve_embedding(&config.read(), &embedding_model_id)?; + let chunk_size = match chunk_size { Some(value) => { println!("Set chunk size: {value}"); @@ -138,9 +140,9 @@ impl Rag { } None => { if *IS_STDOUT_TERMINAL { - set_chunk_size(&model)? + set_chunk_size(&embedding_model)? } else { - let value = model.default_chunk_size(); + let value = embedding_model.default_chunk_size(); println!("Set chunk size: {value}"); value } @@ -161,7 +163,8 @@ impl Rag { } } }; - Ok((model, chunk_size, chunk_overlap)) + + Ok((embedding_model, chunk_size, chunk_overlap)) } pub fn save(&self, path: &Path) -> Result<()> { @@ -176,7 +179,7 @@ impl Rag { let files: Vec<_> = self.data.files.iter().map(|v| &v.path).collect(); let data = json!({ "path": self.path, - "model": self.model.id(), + "embedding_model": self.embedding_model.id(), "chunk_size": self.data.chunk_size, "chunk_overlap": self.data.chunk_overlap, "files": files, @@ -198,13 +201,14 @@ impl Rag { &self, text: &str, top_k: usize, - min_score_vector: f32, - min_score_text: f32, + min_score_vector_search: f32, + min_score_fulltext_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, min_score_text) => { + ret = self.hybird_search(text, top_k, min_score_vector_search, min_score_fulltext_search, rerank) => { ret } _ = watch_abort_signal(abort_signal) => { @@ -290,6 +294,7 @@ impl Rag { self.data.add(rag_files, vector_ids, embeddings); progress(&progress_tx, "Building vector store".into()); self.hnsw = self.data.build_hnsw(); + self.bm25 = self.data.build_bm25(); Ok(()) } @@ -298,22 +303,58 @@ impl Rag { &self, query: &str, top_k: usize, - min_score_vector: f32, - min_score_text: f32, + min_score_vector_search: f32, + min_score_fulltext_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), - self.text_search(query, top_k, min_score_text) + self.vector_search(query, top_k, min_score_vector_search), + self.fulltext_search(query, top_k, min_score_fulltext_search) ); 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 + let fulltext_search_ids = text_search_result?; + debug!("vector_search_ids: {vector_search_ids:?}, fulltext_search_ids: {fulltext_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] + .concat() + .into_iter() + .collect(); + let mut documents = vec![]; + let mut documents_ids = vec![]; + for id in ids { + if let Some(document) = self.data.get(id) { + documents_ids.push(id); + documents.push(document.page_content.to_string()); + } + } + let data = RerankData::new(query.to_string(), documents, top_k); + let list = client.rerank(data).await?; + let ids = list + .into_iter() + .filter_map(|item| { + if item.relevance_score < min_score { + None + } else { + documents_ids.get(item.index).cloned() + } + }) + .collect(); + debug!("rerank_ids: {ids:?}"); + ids + } + None => { + let ids = + reciprocal_rank_fusion(vector_search_ids, fulltext_search_ids, 1.0, 1.0, top_k); + debug!("rrf_ids: {ids:?}"); + ids + } + }; + let output = 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)?; + let document = self.data.get(id)?; Some(document.page_content.clone()) }) .collect(); @@ -352,7 +393,7 @@ impl Rag { Ok(output) } - async fn text_search( + async fn fulltext_search( &self, query: &str, top_k: usize, @@ -369,7 +410,7 @@ impl Rag { ) -> Result<EmbeddingsOutput> { let EmbeddingsData { texts, query } = data; let mut output = vec![]; - let chunks = texts.chunks(self.model.max_concurrent_chunks()); + let chunks = texts.chunks(self.embedding_model.max_concurrent_chunks()); let chunks_len = chunks.len(); progress( &progress_tx, @@ -381,7 +422,7 @@ impl Rag { query, }; let chunk_output = self - .client + .embedding_client .embeddings(chunk_data) .await .context("Failed to create embedding")?; @@ -397,7 +438,7 @@ impl Rag { #[derive(Debug, Clone, Serialize, Deserialize)] pub struct RagData { - pub model: String, + pub embedding_model: String, pub chunk_size: usize, pub chunk_overlap: usize, pub files: Vec<RagFile>, @@ -405,9 +446,9 @@ pub struct RagData { } impl RagData { - pub fn new(model: &str, chunk_size: usize, chunk_overlap: usize) -> Self { + pub fn new(embedding_model: String, chunk_size: usize, chunk_overlap: usize) -> Self { Self { - model: model.to_string(), + embedding_model, chunk_size, chunk_overlap, files: Default::default(), @@ -415,6 +456,13 @@ impl RagData { } } + pub fn get(&self, id: VectorID) -> Option<&RagDocument> { + let (file_index, document_index) = split_vector_id(id); + let file = self.files.get(file_index)?; + let document = file.documents.get(document_index)?; + Some(document) + } + pub fn add( &mut self, files: Vec<RagFile>, |
