Spaces:
Running
Running
File size: 4,204 Bytes
179a518 49f6272 179a518 49f6272 179a518 49f6272 179a518 49f6272 179a518 49f6272 179a518 49f6272 179a518 49f6272 179a518 49f6272 78c73e4 49f6272 179a518 49f6272 179a518 49f6272 179a518 49f6272 179a518 49f6272 179a518 49f6272 179a518 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 | 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,
)
|