RAG_Chatbot / backend /hf_embedding.py
senlinyy's picture
feat: completed initial dev
4b81334
Raw
History Blame Contribute Delete
7.39 kB
from __future__ import annotations
import logging
import os
import time
from typing import Any, Dict, List, Optional, Union
from huggingface_hub import InferenceClient
from huggingface_hub.errors import HfHubHTTPError, InferenceTimeoutError
from llama_index.core.base.embeddings.base import BaseEmbedding, Embedding
from llama_index.core.bridge.pydantic import Field, PrivateAttr
from llama_index.embeddings.huggingface_api.pooling import Pooling
from llama_index.utils.huggingface import format_query, format_text
logger = logging.getLogger(__name__)
DEFAULT_MAX_CHARS = int(os.environ.get("HF_EMBED_MAX_CHARS", "3000"))
DEFAULT_MIN_CHARS = int(os.environ.get("HF_EMBED_MIN_CHARS", "750"))
DEFAULT_MAX_RETRIES = int(os.environ.get("HF_EMBED_RETRIES", "3"))
DEFAULT_BACKOFF_SECONDS = float(os.environ.get("HF_EMBED_RETRY_BACKOFF", "1.0"))
class SyncHuggingFaceInferenceEmbedding(BaseEmbedding):
"""Sync-only embedding adapter for HF Inference API.
The upstream LlamaIndex HF wrapper uses AsyncInferenceClient internally even
for sync calls, which is brittle under uvicorn/asyncio. This adapter uses
only the regular InferenceClient, so indexing and retrieval can run safely
inside a worker thread.
"""
pooling: Optional[Pooling] = Field(default=Pooling.CLS)
query_instruction: Optional[str] = Field(default=None)
text_instruction: Optional[str] = Field(default=None)
model_name: str = Field(default="BAAI/bge-small-en-v1.5")
token: Union[str, bool, None] = Field(default=None)
timeout: Optional[float] = Field(default=None)
headers: Optional[Dict[str, str]] = Field(default=None)
cookies: Optional[Dict[str, str]] = Field(default=None)
max_chars_per_request: int = Field(default=DEFAULT_MAX_CHARS, gt=0)
min_chars_per_request: int = Field(default=DEFAULT_MIN_CHARS, gt=0)
max_retries: int = Field(default=DEFAULT_MAX_RETRIES, ge=1)
retry_backoff_seconds: float = Field(default=DEFAULT_BACKOFF_SECONDS, ge=0.0)
_client: InferenceClient = PrivateAttr()
def __init__(self, **kwargs: Any) -> None:
super().__init__(**kwargs)
self._client = InferenceClient(
model=self.model_name,
token=self.token,
timeout=self.timeout,
headers=self.headers,
cookies=self.cookies,
)
@classmethod
def class_name(cls) -> str:
return "SyncHuggingFaceInferenceEmbedding"
@staticmethod
def _mean_pool_vectors(vectors: List[Embedding]) -> Embedding:
if not vectors:
raise ValueError("Cannot average an empty list of embeddings.")
if len(vectors) == 1:
return vectors[0]
return [sum(values) / len(values) for values in zip(*vectors)]
@staticmethod
def _split_text(text: str, max_chars: int) -> List[str]:
stripped = text.strip()
if len(stripped) <= max_chars:
return [stripped] if stripped else [" "]
paragraphs = [part.strip() for part in stripped.split("\n\n") if part.strip()]
if not paragraphs:
paragraphs = [stripped]
segments: List[str] = []
current = ""
for paragraph in paragraphs:
pieces = [paragraph[i : i + max_chars] for i in range(0, len(paragraph), max_chars)]
for piece in pieces:
candidate = piece if not current else f"{current}\n\n{piece}"
if len(candidate) <= max_chars:
current = candidate
else:
if current:
segments.append(current)
current = piece
if current:
segments.append(current)
return segments or [" "]
def _embed_request(self, text: str) -> Embedding:
embedding = self._client.feature_extraction(
text,
truncate=True,
truncation_direction="right",
)
if len(embedding.shape) == 1:
return embedding.tolist()
embedding = embedding.squeeze(axis=0)
if len(embedding.shape) == 1:
return embedding.tolist()
if self.pooling is None:
raise ValueError(
f"Pooling is required for {self.model_name} because it returned "
"a > 1-D value."
)
return self.pooling(embedding).tolist()
def _embed_with_retry(self, text: str) -> Embedding:
last_error: Exception | None = None
for attempt in range(1, self.max_retries + 1):
try:
return self._embed_request(text)
except (HfHubHTTPError, InferenceTimeoutError) as exc:
last_error = exc
status_code = getattr(getattr(exc, "response", None), "status_code", None)
retryable = isinstance(exc, InferenceTimeoutError) or status_code in (429, 500, 502, 503, 504)
if not retryable or attempt == self.max_retries:
break
delay = self.retry_backoff_seconds * (2 ** (attempt - 1))
logger.warning(
"HF embedding request failed for %s (status=%s, attempt %d/%d). Retrying in %.1fs.",
self.model_name,
status_code,
attempt,
self.max_retries,
delay,
)
time.sleep(delay)
if last_error is not None:
raise last_error
raise RuntimeError("Embedding request failed without an exception.")
def _embed_single(self, text: str, *, max_chars: Optional[int] = None) -> Embedding:
stripped = text.strip() or " "
current_max_chars = max_chars or self.max_chars_per_request
segments = self._split_text(stripped, current_max_chars)
if len(segments) > 1:
embeddings = [self._embed_single(segment, max_chars=current_max_chars) for segment in segments]
return self._mean_pool_vectors(embeddings)
try:
return self._embed_with_retry(segments[0])
except (HfHubHTTPError, InferenceTimeoutError) as exc:
if current_max_chars <= self.min_chars_per_request or len(stripped) <= self.min_chars_per_request:
raise exc
smaller_max_chars = max(self.min_chars_per_request, current_max_chars // 2)
if smaller_max_chars >= current_max_chars:
raise exc
logger.warning(
"HF embedding request for %s failed after retries. Falling back to smaller text segments (%d -> %d chars).",
self.model_name,
current_max_chars,
smaller_max_chars,
)
return self._embed_single(stripped, max_chars=smaller_max_chars)
def _get_query_embedding(self, query: str) -> Embedding:
return self._embed_single(
format_query(query, self.model_name, self.query_instruction)
)
async def _aget_query_embedding(self, query: str) -> Embedding:
return self._get_query_embedding(query)
def _get_text_embedding(self, text: str) -> Embedding:
return self._embed_single(
format_text(text, self.model_name, self.text_instruction)
)
async def _aget_text_embedding(self, text: str) -> Embedding:
return self._get_text_embedding(text)