chatbot / scripts /vector_rag_manager.py
Moderator404's picture
Upload 38 files
a5778f2 verified
Raw
History Blame
4.98 kB
import os
import logging
from pathlib import Path
from typing import List, Dict
import chromadb
from chromadb.config import Settings
from sentence_transformers import SentenceTransformer
from .knowledge_base_manager import KnowledgeBaseManager
os.environ['ANONYMIZED_TELEMETRY'] = 'False'
os.environ['CHROMA_TELEMETRY_ENABLED'] = 'false'
logger = logging.getLogger(__name__)
class VectorRAGManager:
def __init__(self, knowledge_base_path: str, vector_db_path: str = "vector_db", embedding_model: str = "all-MiniLM-L6-v2"):
self.kb_path = Path(knowledge_base_path)
self.vector_db_path = Path(vector_db_path)
self.vector_db_path.mkdir(exist_ok=True, parents=True)
self.kb_manager = KnowledgeBaseManager(knowledge_base_path)
self.chroma_client = chromadb.PersistentClient(path=str(self.vector_db_path), settings=Settings(anonymized_telemetry=False))
self.embedding_model = SentenceTransformer(embedding_model)
self.collection_name = "tmc_documents"
self._init_collection()
self.chunk_size = 800
self.chunk_overlap = 120
def _init_collection(self):
try:
self.collection = self.chroma_client.get_collection(self.collection_name)
except:
self.collection = self.chroma_client.create_collection(self.collection_name)
def chunk_text(self, text: str) -> List[str]:
if len(text) <= self.chunk_size:
return [text]
chunks = []
start = 0
while start < len(text):
end = min(start + self.chunk_size, len(text))
if end < len(text):
for i in range(end, max(start, end-200), -1):
if text[i] in '.!?\n':
end = i+1
break
chunks.append(text[start:end].strip())
start = end - self.chunk_overlap
return [c for c in chunks if len(c) > 50]
def generate_embedding(self, text: str) -> List[float]:
return self.embedding_model.encode(text).tolist()
def index_documents(self, force_reindex: bool = False) -> Dict:
if force_reindex and self.collection.count() > 0:
self.chroma_client.delete_collection(self.collection_name)
self._init_collection()
self.kb_manager.scan_documents()
stats = {"documents": 0, "chunks": 0}
for category, cat_info in self.kb_manager.documents.items():
if category == 'development':
logger.info(f"Skipping development folder")
continue
for doc in cat_info['documents']:
content = self.kb_manager.load_document_content(doc['path'])
if not content:
continue
chunks = self.chunk_text(content)
if not chunks:
continue
ids = []
embeds = []
texts = []
metas = []
for i, chunk in enumerate(chunks):
cid = f"{doc['filename']}_{i}"
ids.append(cid)
texts.append(chunk)
embeds.append(self.generate_embedding(chunk))
metas.append({
"document_title": doc['title'],
"document_path": doc['path'],
"category": category,
"chunk_index": i
})
self.collection.add(ids=ids, embeddings=embeds, documents=texts, metadatas=metas)
stats["documents"] += 1
stats["chunks"] += len(chunks)
logger.info(f"Indexed {stats['documents']} docs, {stats['chunks']} chunks")
return stats
def semantic_search(self, query: str, n_results: int = 5) -> List[Dict]:
if self.collection.count() == 0:
return []
q_emb = self.generate_embedding(query)
results = self.collection.query(query_embeddings=[q_emb], n_results=n_results, include=["documents","metadatas","distances"])
formatted = []
if results['documents'] and results['documents'][0]:
for doc, meta, dist in zip(results['documents'][0], results['metadatas'][0], results['distances'][0]):
similarity = max(0.0, 1.0 - (dist / 2.0))
formatted.append({"document": doc, "metadata": meta, "similarity": similarity})
return formatted
def retrieve_and_rerank_filtered(self, query: str, target_categories: List[str], initial_k: int = 20, final_k: int = 3, similarity_threshold: float = 0.30) -> List[Dict]:
candidates = self.semantic_search(query, n_results=initial_k)
filtered = [c for c in candidates if c['similarity'] >= similarity_threshold and c['metadata'].get('category') in target_categories]
return filtered[:final_k]
def get_collection_stats(self) -> Dict:
return {"total_chunks": self.collection.count()}