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)