Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
localai.py348 linesDownload Raw Back to embeddings
1from __future__ import annotations2 3import logging4import warnings5from typing import (6    Any,7    Callable,8    Dict,9    List,10    Literal,11    Optional,12    Sequence,13    Set,14    Tuple,15    Union,16)17 18from langchain_core.embeddings import Embeddings19from langchain_core.utils import (20    get_from_dict_or_env,21    get_pydantic_field_names,22    pre_init,23)24from pydantic import BaseModel, ConfigDict, Field, model_validator25from tenacity import (26    AsyncRetrying,27    before_sleep_log,28    retry,29    retry_if_exception_type,30    stop_after_attempt,31    wait_exponential,32)33 34logger = logging.getLogger(__name__)35 36 37def _create_retry_decorator(embeddings: LocalAIEmbeddings) -> Callable[[Any], Any]:38    import openai39 40    min_seconds = 441    max_seconds = 1042    # Wait 2^x * 1 second between each retry starting with43    # 4 seconds, then up to 10 seconds, then 10 seconds afterwards44    return retry(45        reraise=True,46        stop=stop_after_attempt(embeddings.max_retries),47        wait=wait_exponential(multiplier=1, min=min_seconds, max=max_seconds),48        retry=(49            retry_if_exception_type(openai.error.Timeout)50            | retry_if_exception_type(openai.error.APIError)51            | retry_if_exception_type(openai.error.APIConnectionError)52            | retry_if_exception_type(openai.error.RateLimitError)53            | retry_if_exception_type(openai.error.ServiceUnavailableError)54        ),55        before_sleep=before_sleep_log(logger, logging.WARNING),56    )57 58 59def _async_retry_decorator(embeddings: LocalAIEmbeddings) -> Any:60    import openai61 62    min_seconds = 463    max_seconds = 1064    # Wait 2^x * 1 second between each retry starting with65    # 4 seconds, then up to 10 seconds, then 10 seconds afterwards66    async_retrying = AsyncRetrying(67        reraise=True,68        stop=stop_after_attempt(embeddings.max_retries),69        wait=wait_exponential(multiplier=1, min=min_seconds, max=max_seconds),70        retry=(71            retry_if_exception_type(openai.error.Timeout)72            | retry_if_exception_type(openai.error.APIError)73            | retry_if_exception_type(openai.error.APIConnectionError)74            | retry_if_exception_type(openai.error.RateLimitError)75            | retry_if_exception_type(openai.error.ServiceUnavailableError)76        ),77        before_sleep=before_sleep_log(logger, logging.WARNING),78    )79 80    def wrap(func: Callable) -> Callable:81        async def wrapped_f(*args: Any, **kwargs: Any) -> Callable:82            async for _ in async_retrying:83                return await func(*args, **kwargs)84            raise AssertionError("this is unreachable")85 86        return wrapped_f87 88    return wrap89 90 91# https://stackoverflow.com/questions/76469415/getting-embeddings-of-length-1-from-langchain-openaiembeddings92def _check_response(response: dict) -> dict:93    if any(len(d["embedding"]) == 1 for d in response["data"]):94        import openai95 96        raise openai.error.APIError("LocalAI API returned an empty embedding")97    return response98 99 100def embed_with_retry(embeddings: LocalAIEmbeddings, **kwargs: Any) -> Any:101    """Use tenacity to retry the embedding call."""102    retry_decorator = _create_retry_decorator(embeddings)103 104    @retry_decorator105    def _embed_with_retry(**kwargs: Any) -> Any:106        response = embeddings.client.create(**kwargs)107        return _check_response(response)108 109    return _embed_with_retry(**kwargs)110 111 112async def async_embed_with_retry(embeddings: LocalAIEmbeddings, **kwargs: Any) -> Any:113    """Use tenacity to retry the embedding call."""114 115    @_async_retry_decorator(embeddings)116    async def _async_embed_with_retry(**kwargs: Any) -> Any:117        response = await embeddings.client.acreate(**kwargs)118        return _check_response(response)119 120    return await _async_embed_with_retry(**kwargs)121 122 123class LocalAIEmbeddings(BaseModel, Embeddings):124    """LocalAI embedding models.125 126    Since LocalAI and OpenAI have 1:1 compatibility between APIs, this class127    uses the ``openai`` Python package's ``openai.Embedding`` as its client.128    Thus, you should have the ``openai`` python package installed, and defeat129    the environment variable ``OPENAI_API_KEY`` by setting to a random string.130    You also need to specify ``OPENAI_API_BASE`` to point to your LocalAI131    service endpoint.132 133    Example:134        .. code-block:: python135 136            from langchain_community.embeddings import LocalAIEmbeddings137            openai = LocalAIEmbeddings(138                openai_api_key="random-string",139                openai_api_base="http://localhost:8080"140            )141 142    """143 144    client: Any = None  #: :meta private:145    model: str = "text-embedding-ada-002"146    deployment: str = model147    openai_api_version: Optional[str] = None148    openai_api_base: Optional[str] = None149    # to support explicit proxy for LocalAI150    openai_proxy: Optional[str] = None151    embedding_ctx_length: int = 8191152    """The maximum number of tokens to embed at once."""153    openai_api_key: Optional[str] = None154    openai_organization: Optional[str] = None155    allowed_special: Union[Literal["all"], Set[str]] = set()156    disallowed_special: Union[Literal["all"], Set[str], Sequence[str]] = "all"157    chunk_size: int = 1000158    """Maximum number of texts to embed in each batch"""159    max_retries: int = 6160    """Maximum number of retries to make when generating."""161    request_timeout: Optional[Union[float, Tuple[float, float]]] = None162    """Timeout in seconds for the LocalAI request."""163    headers: Any = None164    show_progress_bar: bool = False165    """Whether to show a progress bar when embedding."""166    model_kwargs: Dict[str, Any] = Field(default_factory=dict)167    """Holds any model parameters valid for `create` call not explicitly specified."""168 169    model_config = ConfigDict(extra="forbid", protected_namespaces=())170 171    @model_validator(mode="before")172    @classmethod173    def build_extra(cls, values: Dict[str, Any]) -> Any:174        """Build extra kwargs from additional params that were passed in."""175        all_required_field_names = get_pydantic_field_names(cls)176        extra = values.get("model_kwargs", {})177        for field_name in list(values):178            if field_name in extra:179                raise ValueError(f"Found {field_name} supplied twice.")180            if field_name not in all_required_field_names:181                warnings.warn(182                    f"""WARNING! {field_name} is not default parameter.183                    {field_name} was transferred to model_kwargs.184                    Please confirm that {field_name} is what you intended."""185                )186                extra[field_name] = values.pop(field_name)187 188        invalid_model_kwargs = all_required_field_names.intersection(extra.keys())189        if invalid_model_kwargs:190            raise ValueError(191                f"Parameters {invalid_model_kwargs} should be specified explicitly. "192                f"Instead they were passed in as part of `model_kwargs` parameter."193            )194 195        values["model_kwargs"] = extra196        return values197 198    @pre_init199    def validate_environment(cls, values: Dict) -> Dict:200        """Validate that api key and python package exists in environment."""201        values["openai_api_key"] = get_from_dict_or_env(202            values, "openai_api_key", "OPENAI_API_KEY"203        )204        values["openai_api_base"] = get_from_dict_or_env(205            values,206            "openai_api_base",207            "OPENAI_API_BASE",208            default="",209        )210        values["openai_proxy"] = get_from_dict_or_env(211            values,212            "openai_proxy",213            "OPENAI_PROXY",214            default="",215        )216 217        default_api_version = ""218        values["openai_api_version"] = get_from_dict_or_env(219            values,220            "openai_api_version",221            "OPENAI_API_VERSION",222            default=default_api_version,223        )224        values["openai_organization"] = get_from_dict_or_env(225            values,226            "openai_organization",227            "OPENAI_ORGANIZATION",228            default="",229        )230        try:231            import openai232 233            values["client"] = openai.Embedding234        except ImportError:235            raise ImportError(236                "Could not import openai python package. "237                "Please install it with `pip install openai`."238            )239        return values240 241    @property242    def _invocation_params(self) -> Dict:243        openai_args = {244            "model": self.model,245            "request_timeout": self.request_timeout,246            "headers": self.headers,247            "api_key": self.openai_api_key,248            "organization": self.openai_organization,249            "api_base": self.openai_api_base,250            "api_version": self.openai_api_version,251            **self.model_kwargs,252        }253        if self.openai_proxy:254            import openai255 256            openai.proxy = {257                "http": self.openai_proxy,258                "https": self.openai_proxy,259            }260        return openai_args261 262    def _embedding_func(self, text: str, *, engine: str) -> List[float]:263        """Call out to LocalAI's embedding endpoint."""264        # handle large input text265        if self.model.endswith("001"):266            # See: https://github.com/openai/openai-python/issues/418#issuecomment-1525939500267            # replace newlines, which can negatively affect performance.268            text = text.replace("\n", " ")269        return embed_with_retry(270            self,271            input=[text],272            **self._invocation_params,273        )["data"][0]["embedding"]274 275    async def _aembedding_func(self, text: str, *, engine: str) -> List[float]:276        """Call out to LocalAI's embedding endpoint."""277        # handle large input text278        if self.model.endswith("001"):279            # See: https://github.com/openai/openai-python/issues/418#issuecomment-1525939500280            # replace newlines, which can negatively affect performance.281            text = text.replace("\n", " ")282        return (283            await async_embed_with_retry(284                self,285                input=[text],286                **self._invocation_params,287            )288        )["data"][0]["embedding"]289 290    def embed_documents(291        self, texts: List[str], chunk_size: Optional[int] = 0292    ) -> List[List[float]]:293        """Call out to LocalAI's embedding endpoint for embedding search docs.294 295        Args:296            texts: The list of texts to embed.297            chunk_size: The chunk size of embeddings. If None, will use the chunk size298                specified by the class.299 300        Returns:301            List of embeddings, one for each text.302        """303        # call _embedding_func for each text304        return [self._embedding_func(text, engine=self.deployment) for text in texts]305 306    async def aembed_documents(307        self, texts: List[str], chunk_size: Optional[int] = 0308    ) -> List[List[float]]:309        """Call out to LocalAI's embedding endpoint async for embedding search docs.310 311        Args:312            texts: The list of texts to embed.313            chunk_size: The chunk size of embeddings. If None, will use the chunk size314                specified by the class.315 316        Returns:317            List of embeddings, one for each text.318        """319        embeddings = []320        for text in texts:321            response = await self._aembedding_func(text, engine=self.deployment)322            embeddings.append(response)323        return embeddings324 325    def embed_query(self, text: str) -> List[float]:326        """Call out to LocalAI's embedding endpoint for embedding query text.327 328        Args:329            text: The text to embed.330 331        Returns:332            Embedding for the text.333        """334        embedding = self._embedding_func(text, engine=self.deployment)335        return embedding336 337    async def aembed_query(self, text: str) -> List[float]:338        """Call out to LocalAI's embedding endpoint async for embedding query text.339 340        Args:341            text: The text to embed.342 343        Returns:344            Embedding for the text.345        """346        embedding = await self._aembedding_func(text, engine=self.deployment)347        return embedding348