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

View file

@ -7,7 +7,7 @@ from langchain import PromptTemplate
from langchain.agents import Agent from langchain.agents import Agent
from langchain.chains.base import Chain from langchain.chains.base import Chain
from langchain.chat_models.base import BaseChatModel from langchain.chat_models.base import BaseChatModel
from langchain.llms.base import BaseLLM from langchain.base_language import BaseLanguageModel
from langchain.tools import BaseTool from langchain.tools import BaseTool
@ -98,7 +98,7 @@ def import_agent(agent: str) -> Agent:
return import_class(f"langchain.agents.{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""" """Import llm from llm name"""
return import_class(f"langchain.llms.{llm}") return import_class(f"langchain.llms.{llm}")
@ -106,9 +106,8 @@ def import_llm(llm: str) -> BaseLLM:
def import_tool(tool: str) -> BaseTool: def import_tool(tool: str) -> BaseTool:
"""Import tool from tool name""" """Import tool from tool name"""
from langflow.interface.tools.base import tool_creator 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 tool_creator.type_to_loader_dict[tool]["fcn"]
return import_class(f"langchain.tools.{tool}") 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.agents.tools import Tool
from langchain.callbacks.base import BaseCallbackManager from langchain.callbacks.base import BaseCallbackManager
from langchain.chains.loading import load_chain_from_config 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 langchain.llms.loading import load_llm_from_config
from pydantic import ValidationError from pydantic import ValidationError
@ -186,7 +186,7 @@ def load_langchain_type_from_config(config: Dict[str, Any]):
def load_agent_executor_from_config( def load_agent_executor_from_config(
config: dict, config: dict,
llm: Optional[BaseLLM] = None, llm: Optional[BaseLanguageModel] = None,
tools: Optional[list[Tool]] = None, tools: Optional[list[Tool]] = None,
callback_manager: Optional[BaseCallbackManager] = None, callback_manager: Optional[BaseCallbackManager] = None,
**kwargs: Any, **kwargs: Any,

View file

@ -29,7 +29,9 @@ TOOL_INPUTS = {
placeholder="", placeholder="",
value="", 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( "func": TemplateField(
field_type="function", field_type="function",
required=True, required=True,