This commit is contained in:
76
backend/app/core/embeddings.py
Normal file
76
backend/app/core/embeddings.py
Normal file
@@ -0,0 +1,76 @@
|
||||
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
|
||||
Reference in New Issue
Block a user