Spaces:
Runtime error
Runtime error
File size: 4,570 Bytes
aad7814 | 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 138 139 | """Embedding abstraction for the v2 backend.
Default backend is a free, local sentence-transformers model (MiniLM). Setting
``embedding_provider='openai'`` switches to OpenAI embeddings. The active
embedder exposes ``embed_dim`` so the FAISS store can size its index, and a
deterministic fake backend is used when neither dependency is available (tests).
"""
from __future__ import annotations
import hashlib
import logging
import struct
from backend.config import settings
logger = logging.getLogger(__name__)
class Embedder:
"""Minimal embedding interface used by the RAG store."""
embed_dim: int
def embed_documents(self, texts: list[str]) -> list[list[float]]: # pragma: no cover - interface
raise NotImplementedError
def embed_query(self, text: str) -> list[float]: # pragma: no cover - interface
raise NotImplementedError
class LocalEmbedder(Embedder):
"""sentence-transformers backend (default)."""
def __init__(self, model_name: str) -> None:
from sentence_transformers import SentenceTransformer
self._model = SentenceTransformer(model_name)
self.embed_dim = int(self._model.get_sentence_embedding_dimension())
logger.info("LocalEmbedder loaded %s (dim=%d)", model_name, self.embed_dim)
def embed_documents(self, texts: list[str]) -> list[list[float]]:
vecs = self._model.encode(
list(texts), normalize_embeddings=True, convert_to_numpy=True
)
return [v.tolist() for v in vecs]
def embed_query(self, text: str) -> list[float]:
return self.embed_documents([text])[0]
class OpenAIEmbedder(Embedder):
"""OpenAI embeddings backend."""
_DIMS = {
"text-embedding-3-small": 1536,
"text-embedding-3-large": 3072,
"text-embedding-ada-002": 1536,
}
def __init__(self, model_name: str, api_key: str) -> None:
from openai import OpenAI
self._client = OpenAI(
api_key=api_key,
timeout=float(settings.openai_request_timeout_seconds),
)
self._model = model_name
self.embed_dim = self._DIMS.get(model_name, 1536)
logger.info("OpenAIEmbedder using %s (dim=%d)", model_name, self.embed_dim)
def embed_documents(self, texts: list[str]) -> list[list[float]]:
from backend.llm.openai_client import pipeline_timeout
resp = self._client.embeddings.create(
model=self._model,
input=list(texts),
timeout=pipeline_timeout(),
)
return [d.embedding for d in resp.data]
def embed_query(self, text: str) -> list[float]:
return self.embed_documents([text])[0]
class FakeEmbedder(Embedder):
"""Deterministic hash-based embedder for offline tests."""
def __init__(self, dim: int = 384) -> None:
self.embed_dim = dim
def _vec(self, text: str) -> list[float]:
out: list[float] = []
seed = (text or "").encode("utf-8")
counter = 0
while len(out) < self.embed_dim:
h = hashlib.sha256(seed + struct.pack(">I", counter)).digest()
for i in range(0, len(h), 4):
if len(out) >= self.embed_dim:
break
(val,) = struct.unpack(">I", h[i : i + 4])
out.append((val / 0xFFFFFFFF) * 2.0 - 1.0)
counter += 1
norm = sum(x * x for x in out) ** 0.5 or 1.0
return [x / norm for x in out]
def embed_documents(self, texts: list[str]) -> list[list[float]]:
return [self._vec(t) for t in texts]
def embed_query(self, text: str) -> list[float]:
return self._vec(text)
_instance: Embedder | None = None
def get_embedder() -> Embedder:
"""Return the cached embedder singleton selected by configuration."""
global _instance
if _instance is not None:
return _instance
provider = (settings.embedding_provider or "local").lower()
try:
if provider == "openai" and settings.openai_api_key:
_instance = OpenAIEmbedder(settings.openai_embedding_model, settings.openai_api_key)
else:
_instance = LocalEmbedder(settings.local_embedding_model)
except Exception as exc: # noqa: BLE001 - dependency/model load failures → fake
logger.warning("Embedder init failed (%s); using FakeEmbedder (TEST ONLY).", exc)
_instance = FakeEmbedder()
return _instance
def reset_embedder() -> None:
"""Reset the cached embedder (tests / config reloads)."""
global _instance
_instance = None
|