Questro-RAG / src /core /rag.py
BluoCaroot's picture
Deploy Questro RAG service as Docker Space
8b1db68 verified
Raw
History Blame Contribute Delete
5.41 kB
import faiss
import json
import os
import numpy as np
import torch
from sentence_transformers import SentenceTransformer
import sqlite3
class CrossDomainRAGIndex:
def __init__(self, model_name: str = "multi-qa-mpnet-base-dot-v1"):
device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"Initializing embedding model on: {device.upper()}")
self.model = SentenceTransformer(model_name, device=device)
self.dimension = self.model.get_embedding_dimension()
self.index = faiss.IndexHNSWFlat(self.dimension, 32)
self.metadata_store = []
def build_index(self, unified_records: list):
"""Generates embeddings and builds the FAISS index."""
texts = [rec['embedding_text'] for rec in unified_records]
print(f"Generating embeddings for {len(texts)} items...")
embeddings = self.model.encode(
texts,
convert_to_numpy=True,
show_progress_bar=True,
batch_size=128
)
faiss.normalize_L2(embeddings)
self.index.add(embeddings)
# Strip the redundant embedding_text field to save gigabytes of RAM and disk space
for rec in unified_records:
if 'embedding_text' in rec:
del rec['embedding_text']
self.metadata_store.extend(unified_records)
print(f"Successfully indexed {self.index.ntotal} items.")
def save(self, index_path: str, meta_path: str):
"""Saves the FAISS index and metadata to disk."""
print("Writing FAISS index to disk...")
faiss.write_index(self.index, index_path)
print("Writing metadata to SQLite database...")
if os.path.exists(meta_path):
os.remove(meta_path)
conn = sqlite3.connect(meta_path)
cursor = conn.cursor()
cursor.execute('''
CREATE TABLE metadata (
id INTEGER PRIMARY KEY,
data TEXT
)
''')
for i, item in enumerate(self.metadata_store):
cursor.execute('INSERT INTO metadata (id, data) VALUES (?, ?)', (i, json.dumps(item)))
conn.commit()
conn.close()
print("Save complete!")
def load(self, index_path: str, meta_path: str):
"""Loads the FAISS index and metadata from disk."""
print("Loading FAISS index from disk...")
self.index = faiss.read_index(index_path, faiss.IO_FLAG_MMAP)
print("Connecting to metadata SQLite database...")
self.db_conn = sqlite3.connect(meta_path, check_same_thread=False)
print(f"Loaded index with {self.index.ntotal} items.")
def retrieve(self, query: str, top_k: int = 5, blocked_genres: list = None, allow_adult: bool = False) -> list:
"""Retrieves the top_k most similar items to the user's query, applying local filters."""
# Raw query embeds better — sentence transformers are trained on natural text, not lemmatized input
query_embedding = self.model.encode([query], convert_to_numpy=True)
query_embedding = np.atleast_2d(query_embedding).astype(np.float32)
faiss.normalize_L2(query_embedding)
# Fetch a larger pool to support filtering and domain diversity selection
fetch_k = top_k * 10
distances, indices = self.index.search(query_embedding, fetch_k)
blocked_set = set(g.lower() for g in blocked_genres) if blocked_genres else set()
# Collect all valid candidates in score order
candidates = []
cursor = self.db_conn.cursor()
for dist, idx in zip(distances[0], indices[0]):
if idx == -1:
continue
cursor.execute('SELECT data FROM metadata WHERE id = ?', (int(idx),))
row = cursor.fetchone()
if not row:
continue
data = json.loads(row[0])
is_adult_val = data.get("is_adult", False)
is_adult = is_adult_val.lower() == 'true' if isinstance(is_adult_val, str) else bool(is_adult_val)
if not allow_adult and is_adult:
continue
if blocked_set:
item_themes = [t.strip().lower() for t in data.get('themes', '').split(',')]
if not blocked_set.isdisjoint(item_themes):
continue
candidates.append({"score": float(dist), "data": data})
# Guarantee at least one result per domain type (game/movie) to prevent same-domain clustering
results = []
seen_types = set()
deferred = []
for c in candidates:
item_type = c['data'].get('type', 'unknown')
if item_type not in seen_types:
results.append(c)
seen_types.add(item_type)
else:
deferred.append(c)
# Fill remaining slots with the next best candidates regardless of type
for c in deferred:
if len(results) >= top_k:
break
results.append(c)
results.sort(key=lambda x: x['score'])
return results[:top_k]