fix: improved the huggingface emdeddings component, to handle local inference and serverless inference (#6292)
Update huggingface_inference_api.py
This commit is contained in:
parent
f4715407b8
commit
553e3a0b12
1 changed files with 8 additions and 3 deletions
|
|
@ -2,6 +2,8 @@ from urllib.parse import urlparse
|
||||||
|
|
||||||
import requests
|
import requests
|
||||||
from langchain_community.embeddings.huggingface import HuggingFaceInferenceAPIEmbeddings
|
from langchain_community.embeddings.huggingface import HuggingFaceInferenceAPIEmbeddings
|
||||||
|
|
||||||
|
# Next update: use langchain_huggingface
|
||||||
from pydantic import SecretStr
|
from pydantic import SecretStr
|
||||||
from tenacity import retry, stop_after_attempt, wait_fixed
|
from tenacity import retry, stop_after_attempt, wait_fixed
|
||||||
|
|
||||||
|
|
@ -21,7 +23,7 @@ class HuggingFaceInferenceAPIEmbeddingsComponent(LCEmbeddingsModel):
|
||||||
SecretStrInput(
|
SecretStrInput(
|
||||||
name="api_key",
|
name="api_key",
|
||||||
display_name="API Key",
|
display_name="API Key",
|
||||||
advanced=True,
|
advanced=False,
|
||||||
info="Required for non-local inference endpoints. Local inference does not require an API Key.",
|
info="Required for non-local inference endpoints. Local inference does not require an API Key.",
|
||||||
),
|
),
|
||||||
MessageTextInput(
|
MessageTextInput(
|
||||||
|
|
@ -83,11 +85,14 @@ class HuggingFaceInferenceAPIEmbeddingsComponent(LCEmbeddingsModel):
|
||||||
def build_embeddings(self) -> Embeddings:
|
def build_embeddings(self) -> Embeddings:
|
||||||
api_url = self.get_api_url()
|
api_url = self.get_api_url()
|
||||||
|
|
||||||
is_local_url = api_url.startswith(("http://localhost", "http://127.0.0.1"))
|
is_local_url = (
|
||||||
|
api_url.startswith(("http://localhost", "http://127.0.0.1", "http://0.0.0.0", "http://docker"))
|
||||||
|
or "huggingface.co" not in api_url.lower()
|
||||||
|
)
|
||||||
|
|
||||||
if not self.api_key and is_local_url:
|
if not self.api_key and is_local_url:
|
||||||
self.validate_inference_endpoint(api_url)
|
self.validate_inference_endpoint(api_url)
|
||||||
api_key = SecretStr("DummyAPIKeyForLocalDeployment")
|
api_key = SecretStr("APIKeyForLocalDeployment")
|
||||||
elif not self.api_key:
|
elif not self.api_key:
|
||||||
msg = "API Key is required for non-local inference endpoints"
|
msg = "API Key is required for non-local inference endpoints"
|
||||||
raise ValueError(msg)
|
raise ValueError(msg)
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue