codekingpro/portable-devtools
114k
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 