vk_hackathon/index/sparse.py

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