Refactor ConversationChain and LLMCheckerChain components
This commit is contained in:
parent
67aca6dd36
commit
e2c53f1166
3 changed files with 34 additions and 14 deletions
|
|
@ -1,9 +1,9 @@
|
||||||
from typing import Callable, Optional, Union
|
from typing import Optional
|
||||||
|
|
||||||
from langchain.chains import ConversationChain
|
from langchain.chains import ConversationChain
|
||||||
|
|
||||||
from langflow import CustomComponent
|
from langflow import CustomComponent
|
||||||
from langflow.field_typing import BaseLanguageModel, BaseMemory, Chain, Text
|
from langflow.field_typing import BaseLanguageModel, BaseMemory, Text
|
||||||
|
|
||||||
|
|
||||||
class ConversationChainComponent(CustomComponent):
|
class ConversationChainComponent(CustomComponent):
|
||||||
|
|
@ -26,7 +26,7 @@ class ConversationChainComponent(CustomComponent):
|
||||||
inputs: str,
|
inputs: str,
|
||||||
llm: BaseLanguageModel,
|
llm: BaseLanguageModel,
|
||||||
memory: Optional[BaseMemory] = None,
|
memory: Optional[BaseMemory] = None,
|
||||||
) -> Union[Chain, Callable, Text]:
|
) -> Text:
|
||||||
if memory is None:
|
if memory is None:
|
||||||
chain = ConversationChain(llm=llm)
|
chain = ConversationChain(llm=llm)
|
||||||
else:
|
else:
|
||||||
|
|
|
||||||
|
|
@ -1,14 +1,15 @@
|
||||||
from typing import Callable, Union
|
|
||||||
|
|
||||||
from langchain.chains import LLMCheckerChain
|
from langchain.chains import LLMCheckerChain
|
||||||
|
|
||||||
from langflow import CustomComponent
|
from langflow import CustomComponent
|
||||||
from langflow.field_typing import BaseLanguageModel, Chain
|
from langflow.field_typing import BaseLanguageModel, Text
|
||||||
|
|
||||||
|
|
||||||
class LLMCheckerChainComponent(CustomComponent):
|
class LLMCheckerChainComponent(CustomComponent):
|
||||||
display_name = "LLMCheckerChain"
|
display_name = "LLMCheckerChain"
|
||||||
description = ""
|
description = ""
|
||||||
documentation = "https://python.langchain.com/docs/modules/chains/additional/llm_checker"
|
documentation = (
|
||||||
|
"https://python.langchain.com/docs/modules/chains/additional/llm_checker"
|
||||||
|
)
|
||||||
|
|
||||||
def build_config(self):
|
def build_config(self):
|
||||||
return {
|
return {
|
||||||
|
|
@ -17,6 +18,12 @@ class LLMCheckerChainComponent(CustomComponent):
|
||||||
|
|
||||||
def build(
|
def build(
|
||||||
self,
|
self,
|
||||||
|
inputs: str,
|
||||||
llm: BaseLanguageModel,
|
llm: BaseLanguageModel,
|
||||||
) -> Union[Chain, Callable]:
|
) -> Text:
|
||||||
return LLMCheckerChain.from_llm(llm=llm)
|
|
||||||
|
chain = LLMCheckerChain.from_llm(llm=llm)
|
||||||
|
response = chain.invoke({chain.input_key: inputs})
|
||||||
|
result = response.get(chain.output_key)
|
||||||
|
self.status = result
|
||||||
|
return result
|
||||||
|
|
|
||||||
|
|
@ -1,15 +1,17 @@
|
||||||
from typing import Callable, Optional, Union
|
from typing import Optional
|
||||||
|
|
||||||
from langchain.chains import LLMChain, LLMMathChain
|
from langchain.chains import LLMChain, LLMMathChain
|
||||||
|
|
||||||
from langflow import CustomComponent
|
from langflow import CustomComponent
|
||||||
from langflow.field_typing import BaseLanguageModel, BaseMemory, Chain
|
from langflow.field_typing import BaseLanguageModel, BaseMemory, Text
|
||||||
|
|
||||||
|
|
||||||
class LLMMathChainComponent(CustomComponent):
|
class LLMMathChainComponent(CustomComponent):
|
||||||
display_name = "LLMMathChain"
|
display_name = "LLMMathChain"
|
||||||
description = "Chain that interprets a prompt and executes python code to do math."
|
description = "Chain that interprets a prompt and executes python code to do math."
|
||||||
documentation = "https://python.langchain.com/docs/modules/chains/additional/llm_math"
|
documentation = (
|
||||||
|
"https://python.langchain.com/docs/modules/chains/additional/llm_math"
|
||||||
|
)
|
||||||
|
|
||||||
def build_config(self):
|
def build_config(self):
|
||||||
return {
|
return {
|
||||||
|
|
@ -22,10 +24,21 @@ class LLMMathChainComponent(CustomComponent):
|
||||||
|
|
||||||
def build(
|
def build(
|
||||||
self,
|
self,
|
||||||
|
inputs: Text,
|
||||||
llm: BaseLanguageModel,
|
llm: BaseLanguageModel,
|
||||||
llm_chain: LLMChain,
|
llm_chain: LLMChain,
|
||||||
input_key: str = "question",
|
input_key: str = "question",
|
||||||
output_key: str = "answer",
|
output_key: str = "answer",
|
||||||
memory: Optional[BaseMemory] = None,
|
memory: Optional[BaseMemory] = None,
|
||||||
) -> Union[LLMMathChain, Callable, Chain]:
|
) -> Text:
|
||||||
return LLMMathChain(llm=llm, llm_chain=llm_chain, input_key=input_key, output_key=output_key, memory=memory)
|
chain = LLMMathChain(
|
||||||
|
llm=llm,
|
||||||
|
llm_chain=llm_chain,
|
||||||
|
input_key=input_key,
|
||||||
|
output_key=output_key,
|
||||||
|
memory=memory,
|
||||||
|
)
|
||||||
|
response = chain.invoke({input_key: inputs})
|
||||||
|
result = response.get(output_key)
|
||||||
|
self.status = result
|
||||||
|
return result
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue