import logging import threading import time from dataclasses import dataclass from typing import List, Optional, Sequence, Tuple import torch from sentence_transformers import CrossEncoder, SentenceTransformer logger = logging.getLogger(__name__) class ThreadSafeEmbeddings: """Thread-safe SentenceTransformer wrapper used by FAISS indexing/search.""" def __init__(self, model_name: str, device: str): self.model_name = model_name self.device = device self._lock = threading.RLock() started = time.time() logger.info( "[CUSTOM_TOPIC_QUIZ_GENERATOR] Loading embedding model %s on %s", model_name, device, ) self._model = SentenceTransformer(model_name, device=device) logger.info( "[CUSTOM_TOPIC_QUIZ_GENERATOR] Embedding model loaded in %.2fs", time.time() - started, ) def embed_documents(self, texts: List[str]) -> List[List[float]]: if not texts: return [] with self._lock: vectors = self._model.encode( texts, normalize_embeddings=True, convert_to_numpy=True, show_progress_bar=False, ) return vectors.tolist() def embed_query(self, text: str) -> List[float]: with self._lock: vector = self._model.encode( [text], normalize_embeddings=True, convert_to_numpy=True, show_progress_bar=False, )[0] return vector.tolist() class SharedReranker: """Thread-safe multilingual cross-encoder wrapper.""" def __init__(self, model_name: str): self.model_name = model_name self._lock = threading.RLock() started = time.time() logger.info( "[CUSTOM_TOPIC_QUIZ_GENERATOR] Loading reranker model %s", model_name, ) self._model = CrossEncoder( model_name, trust_remote_code=True, automodel_args={"use_flash_attn": False}, ) logger.info( "[CUSTOM_TOPIC_QUIZ_GENERATOR] Reranker loaded in %.2fs", time.time() - started, ) def rank( self, query: str, items: Sequence[Tuple[object, str]], top_k: int, ) -> List[Tuple[object, float]]: if not items or top_k <= 0: return [] pairs = [[query, text] for _, text in items] with self._lock: scores = self._model.predict(pairs) ranked = [ (item, float(score)) for (item, _), score in zip(items, scores) ] ranked.sort(key=lambda pair: pair[1], reverse=True) return ranked[:top_k] @dataclass(frozen=True) class ModelBundle: embeddings: ThreadSafeEmbeddings reranker: Optional[SharedReranker] embedding_model_name: str reranker_model_name: str embedding_device: str @classmethod def load( cls, *, embedding_model_name: str, reranker_model_name: str, use_gpu_for_embeddings: bool, enable_reranker: bool, ) -> "ModelBundle": device = "cuda" if use_gpu_for_embeddings and torch.cuda.is_available() else "cpu" if use_gpu_for_embeddings and device != "cuda": logger.warning( "[CUSTOM_TOPIC_QUIZ_GENERATOR] CUDA requested but unavailable; using CPU." ) embeddings = ThreadSafeEmbeddings(embedding_model_name, device) reranker: Optional[SharedReranker] = None if enable_reranker: try: reranker = SharedReranker(reranker_model_name) except Exception: logger.exception( "[CUSTOM_TOPIC_QUIZ_GENERATOR] Reranker failed to load; " "dense and lexical retrieval remain available." ) return cls( embeddings=embeddings, reranker=reranker, embedding_model_name=embedding_model_name, reranker_model_name=reranker_model_name, embedding_device=device, )