refactor: Update VertexAIEmbeddingsComponent to use new Inputs/Outputs format

This commit is contained in:
Cezar Vasconcelos 2024-06-19 21:04:46 +00:00
commit 4a201d478c

View file

@ -1,76 +1,101 @@
from typing import List, Optional from typing import List, Optional
from langflow.custom import CustomComponent from langflow.base.models.model import LCModelComponent
from langflow.field_typing import Embeddings from langflow.field_typing import Embeddings
from langflow.io import BoolInput, DictInput, FileInput, FloatInput, IntInput, Output, TextInput
class VertexAIEmbeddingsComponent(CustomComponent): class VertexAIEmbeddingsComponent(LCModelComponent):
display_name = "VertexAI Embeddings" display_name = "VertexAI Embeddings"
description = "Generate embeddings using Google Cloud VertexAI models." description = "Generate embeddings using Google Cloud VertexAI models."
icon = "VertexAI"
def build_config(self): inputs = [
return { FileInput(
"credentials": { name="credentials",
"display_name": "Credentials", display_name="Credentials",
"value": "", value="",
"file_types": [".json"], file_types=["json"], # Removed the dot
"field_type": "file", ),
}, DictInput(
"instance": { name="instance",
"display_name": "instance", display_name="Instance",
"advanced": True, advanced=True,
"field_type": "dict", ),
}, TextInput(
"location": { name="location",
"display_name": "Location", display_name="Location",
"value": "us-central1", value="us-central1",
"advanced": True, advanced=True,
}, ),
"max_output_tokens": {"display_name": "Max Output Tokens", "value": 128}, IntInput(
"max_retries": { name="max_output_tokens",
"display_name": "Max Retries", display_name="Max Output Tokens",
"value": 6, value=128,
"advanced": True, ),
}, IntInput(
"model_name": { name="max_retries",
"display_name": "Model Name", display_name="Max Retries",
"value": "textembedding-gecko", value=6,
}, advanced=True,
"n": {"display_name": "N", "value": 1, "advanced": True}, ),
"project": {"display_name": "Project", "advanced": True}, TextInput(
"request_parallelism": { name="model_name",
"display_name": "Request Parallelism", display_name="Model Name",
"value": 5, value="textembedding-gecko",
"advanced": True, ),
}, IntInput(
"stop": {"display_name": "Stop", "advanced": True}, name="n",
"streaming": { display_name="N",
"display_name": "Streaming", value=1,
"value": False, advanced=True,
"advanced": True, ),
}, TextInput(
"temperature": {"display_name": "Temperature", "value": 0.0}, name="project",
"top_k": {"display_name": "Top K", "value": 40, "advanced": True}, display_name="Project",
"top_p": {"display_name": "Top P", "value": 0.95, "advanced": True}, advanced=True,
} ),
IntInput(
name="request_parallelism",
display_name="Request Parallelism",
value=5,
advanced=True,
),
TextInput(
name="stop",
display_name="Stop",
advanced=True,
),
BoolInput(
name="streaming",
display_name="Streaming",
value=False,
advanced=True,
),
FloatInput(
name="temperature",
display_name="Temperature",
value=0.0,
),
IntInput(
name="top_k",
display_name="Top K",
value=40,
advanced=True,
),
FloatInput(
name="top_p",
display_name="Top P",
value=0.95,
advanced=True,
),
]
def build( outputs = [
self, Output(display_name="Embeddings", name="embeddings", method="build_embeddings"),
instance: Optional[str] = None, ]
credentials: Optional[str] = None,
location: str = "us-central1", def build_embeddings(self) -> Embeddings:
max_output_tokens: int = 128,
max_retries: int = 6,
model_name: str = "textembedding-gecko",
n: int = 1,
project: Optional[str] = None,
request_parallelism: int = 5,
stop: Optional[List[str]] = None,
streaming: bool = False,
temperature: float = 0.0,
top_k: int = 40,
top_p: float = 0.95,
) -> Embeddings:
try: try:
from langchain_google_vertexai import VertexAIEmbeddings from langchain_google_vertexai import VertexAIEmbeddings
except ImportError: except ImportError:
@ -79,18 +104,18 @@ class VertexAIEmbeddingsComponent(CustomComponent):
) )
return VertexAIEmbeddings( return VertexAIEmbeddings(
instance=instance, instance=self.instance,
credentials=credentials, credentials=self.credentials,
location=location, location=self.location,
max_output_tokens=max_output_tokens, max_output_tokens=self.max_output_tokens,
max_retries=max_retries, max_retries=self.max_retries,
model_name=model_name, model_name=self.model_name,
n=n, n=self.n,
project=project, project=self.project,
request_parallelism=request_parallelism, request_parallelism=self.request_parallelism,
stop=stop, stop=self.stop,
streaming=streaming, streaming=self.streaming,
temperature=temperature, temperature=self.temperature,
top_k=top_k, top_k=self.top_k,
top_p=top_p, top_p=self.top_p,
) )