38 lines
1.1 KiB
Python
38 lines
1.1 KiB
Python
"""
|
|
Сервис для вычисления эмбеддингов текстов
|
|
"""
|
|
from functools import lru_cache
|
|
from typing import Iterable
|
|
import numpy as np
|
|
from sentence_transformers import SentenceTransformer
|
|
|
|
|
|
class EmbeddingService:
|
|
|
|
def __init__(self, model_name: str | None = None):
|
|
self.model_name = model_name or "intfloat/multilingual-e5-base"
|
|
self._model = None
|
|
|
|
@property
|
|
def model(self) -> SentenceTransformer:
|
|
if self._model is None:
|
|
self._model = SentenceTransformer(self.model_name)
|
|
return self._model
|
|
|
|
def embed_texts(self, texts: Iterable[str]) -> list[list[float]]:
|
|
embeddings = self.model.encode(
|
|
list(texts),
|
|
batch_size=8,
|
|
show_progress_bar=False,
|
|
normalize_embeddings=True,
|
|
)
|
|
return [np.array(v, dtype=np.float32).tolist() for v in embeddings]
|
|
|
|
def embed_query(self, text: str) -> list[float]:
|
|
return self.embed_texts([text])[0]
|
|
|
|
@lru_cache(maxsize=1)
|
|
def model_version(self) -> str:
|
|
return self.model_name
|
|
|