Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
huggingface_hub.py160 linesDownload Raw Back to embeddings
1import json2from typing import Any, Dict, List, Optional3 4from langchain_core._api import deprecated5from langchain_core.embeddings import Embeddings6from langchain_core.utils import get_from_dict_or_env7from pydantic import BaseModel, ConfigDict, model_validator8from typing_extensions import Self9 10DEFAULT_MODEL = "sentence-transformers/all-mpnet-base-v2"11VALID_TASKS = ("feature-extraction",)12 13 14@deprecated(15    since="0.2.2",16    removal="1.0",17    alternative_import="langchain_huggingface.HuggingFaceEndpointEmbeddings",18)19class HuggingFaceHubEmbeddings(BaseModel, Embeddings):20    """HuggingFaceHub embedding models.21 22    To use, you should have the ``huggingface_hub`` python package installed, and the23    environment variable ``HUGGINGFACEHUB_API_TOKEN`` set with your API token, or pass24    it as a named parameter to the constructor.25 26    Example:27        .. code-block:: python28 29            from langchain_community.embeddings import HuggingFaceHubEmbeddings30            model = "sentence-transformers/all-mpnet-base-v2"31            hf = HuggingFaceHubEmbeddings(32                model=model,33                task="feature-extraction",34                huggingfacehub_api_token="my-api-key",35            )36    """37 38    client: Any = None  #: :meta private:39    async_client: Any = None  #: :meta private:40    model: Optional[str] = None41    """Model name to use."""42    repo_id: Optional[str] = None43    """Huggingfacehub repository id, for backward compatibility."""44    task: Optional[str] = "feature-extraction"45    """Task to call the model with."""46    model_kwargs: Optional[dict] = None47    """Keyword arguments to pass to the model."""48 49    huggingfacehub_api_token: Optional[str] = None50 51    model_config = ConfigDict(extra="forbid", protected_namespaces=())52 53    @model_validator(mode="before")54    @classmethod55    def validate_environment(cls, values: Dict) -> Any:56        """Validate that api key and python package exists in environment."""57        huggingfacehub_api_token = get_from_dict_or_env(58            values, "huggingfacehub_api_token", "HUGGINGFACEHUB_API_TOKEN"59        )60 61        try:62            from huggingface_hub import AsyncInferenceClient, InferenceClient63 64            if values.get("model"):65                values["repo_id"] = values["model"]66            elif values.get("repo_id"):67                values["model"] = values["repo_id"]68            else:69                values["model"] = DEFAULT_MODEL70                values["repo_id"] = DEFAULT_MODEL71 72            client = InferenceClient(73                model=values["model"],74                token=huggingfacehub_api_token,75            )76 77            async_client = AsyncInferenceClient(78                model=values["model"],79                token=huggingfacehub_api_token,80            )81 82            values["client"] = client83            values["async_client"] = async_client84 85        except ImportError:86            raise ImportError(87                "Could not import huggingface_hub python package. "88                "Please install it with `pip install huggingface_hub`."89            )90        return values91 92    @model_validator(mode="after")93    def post_init(self) -> Self:94        """Post init validation for the class."""95        if self.task not in VALID_TASKS:96            raise ValueError(97                f"Got invalid task {self.task}, "98                f"currently only {VALID_TASKS} are supported"99            )100        return self101 102    def embed_documents(self, texts: List[str]) -> List[List[float]]:103        """Call out to HuggingFaceHub's embedding endpoint for embedding search docs.104 105        Args:106            texts: The list of texts to embed.107 108        Returns:109            List of embeddings, one for each text.110        """111        # replace newlines, which can negatively affect performance.112        texts = [text.replace("\n", " ") for text in texts]113        _model_kwargs = self.model_kwargs or {}114        #  api doc: https://huggingface.github.io/text-embeddings-inference/#/Text%20Embeddings%20Inference/embed115        responses = self.client.post(116            json={"inputs": texts, **_model_kwargs}, task=self.task117        )118        return json.loads(responses.decode())119 120    async def aembed_documents(self, texts: List[str]) -> List[List[float]]:121        """Async Call to HuggingFaceHub's embedding endpoint for embedding search docs.122 123        Args:124            texts: The list of texts to embed.125 126        Returns:127            List of embeddings, one for each text.128        """129        # replace newlines, which can negatively affect performance.130        texts = [text.replace("\n", " ") for text in texts]131        _model_kwargs = self.model_kwargs or {}132        responses = await self.async_client.post(133            json={"inputs": texts, "parameters": _model_kwargs}, task=self.task134        )135        return json.loads(responses.decode())136 137    def embed_query(self, text: str) -> List[float]:138        """Call out to HuggingFaceHub's embedding endpoint for embedding query text.139 140        Args:141            text: The text to embed.142 143        Returns:144            Embeddings for the text.145        """146        response = self.embed_documents([text])[0]147        return response148 149    async def aembed_query(self, text: str) -> List[float]:150        """Async Call to HuggingFaceHub's embedding endpoint for embedding query text.151 152        Args:153            text: The text to embed.154 155        Returns:156            Embeddings for the text.157        """158        response = (await self.aembed_documents([text]))[0]159        return response160 
codekingpro/portable-devtools · Team Ai