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
|
||||
from langchain_community.embeddings.huggingface import HuggingFaceInferenceAPIEmbeddings
|
||||
|
||||
# Next update: use langchain_huggingface
|
||||
from pydantic import SecretStr
|
||||
from tenacity import retry, stop_after_attempt, wait_fixed
|
||||
|
||||
|
|
@ -21,7 +23,7 @@ class HuggingFaceInferenceAPIEmbeddingsComponent(LCEmbeddingsModel):
|
|||
SecretStrInput(
|
||||
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.",
|
||||
),
|
||||
MessageTextInput(
|
||||
|
|
@ -83,11 +85,14 @@ class HuggingFaceInferenceAPIEmbeddingsComponent(LCEmbeddingsModel):
|
|||
def build_embeddings(self) -> Embeddings:
|
||||
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:
|
||||
self.validate_inference_endpoint(api_url)
|
||||
api_key = SecretStr("DummyAPIKeyForLocalDeployment")
|
||||
api_key = SecretStr("APIKeyForLocalDeployment")
|
||||
elif not self.api_key:
|
||||
msg = "API Key is required for non-local inference endpoints"
|
||||
raise ValueError(msg)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue