Spaces:
Sleeping
Sleeping
File size: 7,390 Bytes
4b81334 | 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 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 | 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)
|