Files
tset-antiplagiat/backend/app/core/embeddings.py
jze9 2a14350ee3
Some checks failed
Deploy / deploy (push) Has been cancelled
test build
2026-05-18 01:14:40 +05:00

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