import numpy as np import faiss from sentence_transformers import SentenceTransformer from typing import List, Tuple import os from app.config import settings class EmbeddingIndex: def __init__(self, model_name: str = "sentence-transformers/paraphrase-multilingual-mpnet-base-v2"): self.model = SentenceTransformer(model_name) self.embedding_dim = self.model.get_sentence_embedding_dimension() self.index = faiss.IndexHNSWFlat(self.embedding_dim, 32) self.doc_ids = [] self.embeddings = None # Try to load existing index self._load_index() def add_document(self, doc_id: str, text: str): embedding = self.model.encode([text], convert_to_numpy=True)[0] self.index.add(np.array([embedding], dtype=np.float32)) self.doc_ids.append(doc_id) if self.embeddings is None: self.embeddings = np.array([embedding]) else: self.embeddings = np.vstack([self.embeddings, embedding]) def batch_add(self, docs: List[Tuple[str, str]], batch_size: int = 32): doc_ids = [doc[0] for doc in docs] texts = [doc[1] for doc in docs] for i in range(0, len(texts), batch_size): batch_texts = texts[i:i+batch_size] embeddings = self.model.encode(batch_texts, convert_to_numpy=True) self.index.add(embeddings.astype(np.float32)) self.doc_ids.extend(doc_ids[i:i+batch_size]) if self.embeddings is None: self.embeddings = embeddings else: self.embeddings = np.vstack([self.embeddings, embeddings]) def search(self, text: str, top_k: int = 10) -> List[Tuple[str, float]]: query_embedding = self.model.encode([text], convert_to_numpy=True)[0] query_embedding = np.array([query_embedding], dtype=np.float32) distances, indices = self.index.search(query_embedding, top_k) results = [] for idx, distance in zip(indices[0], distances[0]): if idx < len(self.doc_ids): similarity = 1.0 / (1.0 + distance) results.append((self.doc_ids[idx], similarity)) return results def _load_index(self): index_path = os.path.join(settings.INDEX_DIR, "faiss_index") ids_path = os.path.join(settings.INDEX_DIR, "doc_ids.npy") if os.path.exists(index_path) and os.path.exists(ids_path): try: self.index = faiss.read_index(index_path) self.doc_ids = list(np.load(ids_path, allow_pickle=True)) except Exception as e: print(f"Failed to load index: {e}") def save_index(self): os.makedirs(settings.INDEX_DIR, exist_ok=True) faiss.write_index(self.index, os.path.join(settings.INDEX_DIR, "faiss_index")) np.save(os.path.join(settings.INDEX_DIR, "doc_ids.npy"), np.array(self.doc_ids, dtype=object)) def get_index_size(self) -> int: return self.index.ntotal