forked from zovos/vk_hackathon
39 lines
1.1 KiB
Python
39 lines
1.1 KiB
Python
import logging
|
|
import os
|
|
from functools import lru_cache
|
|
|
|
from index_schemas import SparseVector
|
|
|
|
SPARSE_MODEL_NAME = "Qdrant/bm25"
|
|
FASTEMBED_CACHE_PATH = "/models/fastembed"
|
|
MAX_SPARSE_TEXT_CHARS = int(os.getenv("MAX_SPARSE_TEXT_CHARS", "512"))
|
|
|
|
logger = logging.getLogger("index-service")
|
|
|
|
|
|
@lru_cache(maxsize=1)
|
|
def get_sparse_model():
|
|
from fastembed import SparseTextEmbedding
|
|
|
|
logger.info("Loading sparse model %s from cache %s", SPARSE_MODEL_NAME, FASTEMBED_CACHE_PATH)
|
|
return SparseTextEmbedding(model_name=SPARSE_MODEL_NAME)
|
|
|
|
|
|
def _prepare_text(text: str) -> str:
|
|
if not text:
|
|
return ""
|
|
return text[:MAX_SPARSE_TEXT_CHARS]
|
|
|
|
|
|
def embed_sparse_texts(texts: list[str]) -> list[SparseVector]:
|
|
model = get_sparse_model()
|
|
prepared = [_prepare_text(t) for t in texts]
|
|
result: list[SparseVector] = []
|
|
for item in model.embed(prepared):
|
|
result.append(
|
|
SparseVector(
|
|
indices=[int(i) for i in item.indices.tolist()],
|
|
values=[float(v) for v in item.values.tolist()],
|
|
)
|
|
)
|
|
return result
|