Spaces:
Sleeping
Sleeping
| 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, | |
| ) | |
| def class_name(cls) -> str: | |
| return "SyncHuggingFaceInferenceEmbedding" | |
| 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)] | |
| 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) | |