Spaces:
Running
Running
| 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] | |
| class ModelBundle: | |
| embeddings: ThreadSafeEmbeddings | |
| reranker: Optional[SharedReranker] | |
| embedding_model_name: str | |
| reranker_model_name: str | |
| embedding_device: str | |
| 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, | |
| ) | |