77 lines
3.0 KiB
Python
77 lines
3.0 KiB
Python
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
|