refactor: Cohere and Mistral Embeddings, new Inputs/Outputs format

This commit is contained in:
Cezar Vasconcelos 2024-06-19 20:59:12 +00:00
commit 5d21466525
2 changed files with 77 additions and 79 deletions

View file

@ -1,38 +1,44 @@
from typing import Optional
from langchain_community.embeddings.cohere import CohereEmbeddings from langchain_community.embeddings.cohere import CohereEmbeddings
from langflow.custom import CustomComponent from langflow.base.models.model import LCModelComponent
from langflow.field_typing import Embeddings
from langflow.io import BoolInput, DictInput, DropdownInput, FloatInput, IntInput, Output, SecretStrInput, TextInput
class CohereEmbeddingsComponent(CustomComponent): class CohereEmbeddingsComponent(LCModelComponent):
display_name = "Cohere Embeddings" display_name = "Cohere Embeddings"
description = "Generate embeddings using Cohere models." description = "Generate embeddings using Cohere models."
icon = "Cohere"
inputs = [
SecretStrInput(name="cohere_api_key", display_name="Cohere API Key"),
DropdownInput(
name="model",
display_name="Model",
advanced=True,
options=[
"embed-english-v2.0",
"embed-multilingual-v2.0",
"embed-english-light-v2.0",
"embed-multilingual-light-v2.0",
],
value="embed-english-v2.0",
),
TextInput(name="truncate", display_name="Truncate", advanced=True),
IntInput(name="max_retries", display_name="Max Retries", value=3, advanced=True),
TextInput(name="user_agent", display_name="User Agent", advanced=True, value="langchain"),
FloatInput(name="request_timeout", display_name="Request Timeout", advanced=True),
]
def build_config(self): outputs = [
return { Output(display_name="Embeddings", name="embeddings", method="build_embeddings"),
"cohere_api_key": {"display_name": "Cohere API Key", "password": True}, ]
"model": {"display_name": "Model", "default": "embed-english-v2.0", "advanced": True},
"truncate": {"display_name": "Truncate", "advanced": True},
"max_retries": {"display_name": "Max Retries", "advanced": True},
"user_agent": {"display_name": "User Agent", "advanced": True},
"request_timeout": {"display_name": "Request Timeout", "advanced": True},
}
def build( def build_embeddings(self) -> Embeddings:
self, return CohereEmbeddings(
request_timeout: Optional[float] = None, cohere_api_key=self.cohere_api_key,
cohere_api_key: str = "", model=self.model,
max_retries: int = 3, truncate=self.truncate,
model: str = "embed-english-v2.0", max_retries=self.max_retries,
truncate: Optional[str] = None, user_agent=self.user_agent,
user_agent: str = "langchain", request_timeout=self.request_timeout or None,
) -> CohereEmbeddings:
return CohereEmbeddings( # type: ignore
max_retries=max_retries,
user_agent=user_agent,
request_timeout=request_timeout,
cohere_api_key=cohere_api_key,
model=model,
truncate=truncate,
) )

View file

@ -1,64 +1,56 @@
from langchain_mistralai.embeddings import MistralAIEmbeddings from langchain_mistralai.embeddings import MistralAIEmbeddings
from pydantic.v1 import SecretStr from pydantic.v1 import SecretStr
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 DropdownInput, IntInput, Output, SecretStrInput, TextInput
class MistralAIEmbeddingsComponent(CustomComponent): class MistralAIEmbeddingsComponent(LCModelComponent):
display_name = "MistralAI Embeddings" display_name = "MistralAI Embeddings"
description = "Generate embeddings using MistralAI models." description = "Generate embeddings using MistralAI models."
icon = "MistralAI"
def build_config(self): inputs = [
return { DropdownInput(
"model": { name="model",
"display_name": "Model", display_name="Model",
"advanced": False, advanced=False,
"options": ["mistral-embed"], options=["mistral-embed"],
"value": "mistral-embed", value="mistral-embed",
}, ),
"mistral_api_key": { SecretStrInput(name="mistral_api_key", display_name="Mistral API Key"),
"display_name": "Mistral API Key", IntInput(
"password": True, name="max_concurrent_requests",
"advanced": False, display_name="Max Concurrent Requests",
}, advanced=True,
"max_concurrent_requests": { value=64,
"display_name": "Max Concurrent Requests", ),
"advanced": True, IntInput(name="max_retries", display_name="Max Retries", advanced=True, value=5),
"value": 64, IntInput(name="timeout", display_name="Request Timeout", advanced=True, value=120),
}, TextInput(
"max_retries": { name="endpoint",
"display_name": "Max Retries", display_name="API Endpoint",
"advanced": True, advanced=True,
"value": 5, value="https://api.mistral.ai/v1/",
}, ),
"timeout": { ]
"display_name": "Request Timeout",
"advanced": True,
"value": 120,
},
"endpoint": {"display_name": "API Endpoint", "advanced": True, "value": "https://api.mistral.ai/v1/"},
}
def build( outputs = [
self, Output(display_name="Embeddings", name="embeddings", method="build_embeddings"),
mistral_api_key: str, ]
model: str = "mistral-embed",
max_concurrent_requests: int = 64, def build_embeddings(self) -> Embeddings:
max_retries: int = 5, if not self.mistral_api_key:
timeout: int = 120, raise ValueError("Mistral API Key is required")
endpoint: str = "https://api.mistral.ai/v1/",
) -> Embeddings: api_key = SecretStr(self.mistral_api_key)
if mistral_api_key:
api_key = SecretStr(mistral_api_key)
else:
api_key = None
return MistralAIEmbeddings( return MistralAIEmbeddings(
api_key=api_key, api_key=api_key,
model=model, model=self.model,
endpoint=endpoint, endpoint=self.endpoint,
max_concurrent_requests=max_concurrent_requests, max_concurrent_requests=self.max_concurrent_requests,
max_retries=max_retries, max_retries=self.max_retries,
timeout=timeout, timeout=self.timeout,
) )