refactor(langflow): rename BaseLLM to BaseLanguageModel in multiple files

fix(langflow): fix type hints in multiple files
This commit is contained in:
Gabriel Almeida 2023-05-08 15:05:28 -03:00
commit 8810a0392a
4 changed files with 14 additions and 11 deletions

View file

@ -27,7 +27,7 @@ from langchain.agents.agent_toolkits.vectorstore.prompt import (
from langchain.agents.mrkl.prompt import FORMAT_INSTRUCTIONS
from langchain.agents.mrkl.prompt import FORMAT_INSTRUCTIONS as SQL_FORMAT_INSTRUCTIONS
from langchain.base_language import BaseLanguageModel
from langchain.llms.base import BaseLLM
from langchain.memory.chat_memory import BaseChatMemory
from langchain.sql_database import SQLDatabase
from langchain.tools.python.tool import PythonAstREPLTool
@ -134,7 +134,7 @@ class VectorStoreAgent(AgentExecutor):
@classmethod
def from_toolkit_and_llm(
cls, llm: BaseLLM, vectorstoreinfo: VectorStoreInfo, **kwargs: Any
cls, llm: BaseLanguageModel, vectorstoreinfo: VectorStoreInfo, **kwargs: Any
):
"""Construct a vectorstore agent from an LLM and tools."""
@ -171,7 +171,9 @@ class SQLAgent(AgentExecutor):
super().__init__(*args, **kwargs)
@classmethod
def from_toolkit_and_llm(cls, llm: BaseLLM, database_uri: str, **kwargs: Any):
def from_toolkit_and_llm(
cls, llm: BaseLanguageModel, database_uri: str, **kwargs: Any
):
"""Construct a sql agent from an LLM and tools."""
db = SQLDatabase.from_uri(database_uri)
toolkit = SQLDatabaseToolkit(db=db, llm=llm)
@ -275,7 +277,7 @@ class InitializeAgent(AgentExecutor):
@classmethod
def initialize(
cls,
llm: BaseLLM,
llm: BaseLanguageModel,
tools: List[Tool],
agent: str,
memory: Optional[BaseChatMemory] = None,

View file

@ -7,7 +7,7 @@ from langchain import PromptTemplate
from langchain.agents import Agent
from langchain.chains.base import Chain
from langchain.chat_models.base import BaseChatModel
from langchain.llms.base import BaseLLM
from langchain.base_language import BaseLanguageModel
from langchain.tools import BaseTool
@ -98,7 +98,7 @@ def import_agent(agent: str) -> Agent:
return import_class(f"langchain.agents.{agent}")
def import_llm(llm: str) -> BaseLLM:
def import_llm(llm: str) -> BaseLanguageModel:
"""Import llm from llm name"""
return import_class(f"langchain.llms.{llm}")
@ -106,9 +106,8 @@ def import_llm(llm: str) -> BaseLLM:
def import_tool(tool: str) -> BaseTool:
"""Import tool from tool name"""
from langflow.interface.tools.base import tool_creator
from langflow.interface.tools.constants import ALL_TOOLS_NAMES
if tool in ALL_TOOLS_NAMES:
if tool in tool_creator.type_to_loader_dict:
return tool_creator.type_to_loader_dict[tool]["fcn"]
return import_class(f"langchain.tools.{tool}")

View file

@ -15,7 +15,7 @@ from langchain.agents.loading import load_agent_from_config
from langchain.agents.tools import Tool
from langchain.callbacks.base import BaseCallbackManager
from langchain.chains.loading import load_chain_from_config
from langchain.llms.base import BaseLLM
from langchain.base_language import BaseLanguageModel
from langchain.llms.loading import load_llm_from_config
from pydantic import ValidationError
@ -186,7 +186,7 @@ def load_langchain_type_from_config(config: Dict[str, Any]):
def load_agent_executor_from_config(
config: dict,
llm: Optional[BaseLLM] = None,
llm: Optional[BaseLanguageModel] = None,
tools: Optional[list[Tool]] = None,
callback_manager: Optional[BaseCallbackManager] = None,
**kwargs: Any,

View file

@ -29,7 +29,9 @@ TOOL_INPUTS = {
placeholder="",
value="",
),
"llm": TemplateField(field_type="BaseLLM", required=True, is_list=False, show=True),
"llm": TemplateField(
field_type="BaseLanguageModel", required=True, is_list=False, show=True
),
"func": TemplateField(
field_type="function",
required=True,