feat: add ZeroShotPrompt
This commit is contained in:
parent
6a0f5322d9
commit
e909a938c9
5 changed files with 141 additions and 8 deletions
40
langflow/backend/customs.py
Normal file
40
langflow/backend/customs.py
Normal file
|
|
@ -0,0 +1,40 @@
|
||||||
|
from langchain.agents.mrkl import prompt
|
||||||
|
|
||||||
|
|
||||||
|
def get_custom_prompts():
|
||||||
|
return {
|
||||||
|
"ZeroShotPrompt": {
|
||||||
|
"template": {
|
||||||
|
"_type": "zero_shot",
|
||||||
|
"prefix": {
|
||||||
|
"type": "str",
|
||||||
|
"required": False,
|
||||||
|
"placeholder": "",
|
||||||
|
"list": False,
|
||||||
|
"show": True,
|
||||||
|
"multiline": True,
|
||||||
|
"value": prompt.PREFIX,
|
||||||
|
},
|
||||||
|
"suffix": {
|
||||||
|
"type": "str",
|
||||||
|
"required": True,
|
||||||
|
"placeholder": "",
|
||||||
|
"list": False,
|
||||||
|
"show": True,
|
||||||
|
"multiline": True,
|
||||||
|
"value": prompt.SUFFIX,
|
||||||
|
},
|
||||||
|
"format_instructions": {
|
||||||
|
"type": "str",
|
||||||
|
"required": False,
|
||||||
|
"placeholder": "",
|
||||||
|
"list": False,
|
||||||
|
"show": True,
|
||||||
|
"multiline": True,
|
||||||
|
"value": prompt.FORMAT_INSTRUCTIONS,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"description": "Prompt template for Zero Shot Agent.",
|
||||||
|
"base_classes": ["BasePromptTemplate"],
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -9,6 +9,7 @@ from langchain.llms.loading import load_llm_from_config
|
||||||
from langchain.prompts.loading import load_prompt_from_config
|
from langchain.prompts.loading import load_prompt_from_config
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
|
|
||||||
# build router
|
# build router
|
||||||
router = APIRouter()
|
router = APIRouter()
|
||||||
|
|
||||||
|
|
@ -26,6 +27,10 @@ def get_type_list():
|
||||||
|
|
||||||
@router.get("/all")
|
@router.get("/all")
|
||||||
def get_all():
|
def get_all():
|
||||||
|
# library_prompts = {
|
||||||
|
# prompt: signature.get_prompt(prompt) for prompt in list_endpoints.list_prompts()
|
||||||
|
# }
|
||||||
|
# custom_prompts = customs.get_custom_prompts()
|
||||||
return {
|
return {
|
||||||
"chains": {
|
"chains": {
|
||||||
chain: signature.get_chain(chain) for chain in list_endpoints.list_chains()
|
chain: signature.get_chain(chain) for chain in list_endpoints.list_chains()
|
||||||
|
|
@ -33,6 +38,7 @@ def get_all():
|
||||||
"agents": {
|
"agents": {
|
||||||
agent: signature.get_agent(agent) for agent in list_endpoints.list_agents()
|
agent: signature.get_agent(agent) for agent in list_endpoints.list_agents()
|
||||||
},
|
},
|
||||||
|
# "prompts": {**library_prompts, **custom_prompts},
|
||||||
"prompts": {
|
"prompts": {
|
||||||
prompt: signature.get_prompt(prompt)
|
prompt: signature.get_prompt(prompt)
|
||||||
for prompt in list_endpoints.list_prompts()
|
for prompt in list_endpoints.list_prompts()
|
||||||
|
|
@ -67,8 +73,22 @@ def get_all():
|
||||||
|
|
||||||
@router.post("/predict")
|
@router.post("/predict")
|
||||||
def get_load(data: dict[str, Any]):
|
def get_load(data: dict[str, Any]):
|
||||||
|
# Get type list
|
||||||
type_list = get_type_list()
|
type_list = get_type_list()
|
||||||
|
|
||||||
|
# Substitute ZeroShotPromt with PromptTemplate
|
||||||
|
for node in data['nodes']:
|
||||||
|
if node["data"]["type"] == "ZeroShotPrompt":
|
||||||
|
# Build Prompt Template
|
||||||
|
tools = [
|
||||||
|
tool
|
||||||
|
for tool in data['nodes']
|
||||||
|
if tool["type"] != "chatOutputNode"
|
||||||
|
and "Tool" in tool["data"]["node"]["base_classes"]
|
||||||
|
]
|
||||||
|
node["data"] = build_prompt_template(prompt=node["data"], tools=tools)
|
||||||
|
break
|
||||||
|
|
||||||
# Add input variables
|
# Add input variables
|
||||||
data = payload.extract_input_variables(data)
|
data = payload.extract_input_variables(data)
|
||||||
|
|
||||||
|
|
@ -96,12 +116,75 @@ def get_load(data: dict[str, Any]):
|
||||||
else:
|
else:
|
||||||
return {"result": "Error: Type should be either agent, chain or llm"}
|
return {"result": "Error: Type should be either agent, chain or llm"}
|
||||||
|
|
||||||
# elif extracted_json["_type"] in type_list["prompts"]:
|
|
||||||
# loaded = load_prompt_from_config(extracted_json)
|
|
||||||
# print(loaded.format(product=''))
|
|
||||||
|
|
||||||
# return {'result': loaded.format(product=message)}
|
def build_prompt_template(prompt, tools):
|
||||||
|
prefix = prompt["node"]["template"]["prefix"]["value"]
|
||||||
|
suffix = prompt["node"]["template"]["suffix"]["value"]
|
||||||
|
format_instructions = prompt["node"]["template"]["format_instructions"]["value"]
|
||||||
|
|
||||||
# if type in a["prompts"]:
|
tool_strings = "\n".join(
|
||||||
|
[
|
||||||
|
f"{tool['data']['node']['name']}: {tool['data']['node']['description']}"
|
||||||
|
for tool in tools
|
||||||
|
]
|
||||||
|
)
|
||||||
|
tool_names = ", ".join([tool["data"]["node"]["name"] for tool in tools])
|
||||||
|
format_instructions = format_instructions.format(tool_names=tool_names)
|
||||||
|
value = "\n\n".join([prefix, tool_strings, format_instructions, suffix])
|
||||||
|
|
||||||
# return a
|
prompt["type"] = "PromptTemplate"
|
||||||
|
# prompt["value"] = value
|
||||||
|
|
||||||
|
prompt["node"] = {
|
||||||
|
"template": {
|
||||||
|
"_type": "prompt",
|
||||||
|
"input_variables": {
|
||||||
|
"type": "str",
|
||||||
|
"required": True,
|
||||||
|
"placeholder": "",
|
||||||
|
"list": True,
|
||||||
|
"show": False,
|
||||||
|
"multiline": False,
|
||||||
|
},
|
||||||
|
"output_parser": {
|
||||||
|
"type": "BaseOutputParser",
|
||||||
|
"required": False,
|
||||||
|
"placeholder": "",
|
||||||
|
"list": False,
|
||||||
|
"show": False,
|
||||||
|
"multline": False,
|
||||||
|
"value": None,
|
||||||
|
},
|
||||||
|
"template": {
|
||||||
|
"type": "str",
|
||||||
|
"required": True,
|
||||||
|
"placeholder": "",
|
||||||
|
"list": False,
|
||||||
|
"show": True,
|
||||||
|
"multiline": True,
|
||||||
|
"value": value,
|
||||||
|
},
|
||||||
|
"template_format": {
|
||||||
|
"type": "str",
|
||||||
|
"required": False,
|
||||||
|
"placeholder": "",
|
||||||
|
"list": False,
|
||||||
|
"show": False,
|
||||||
|
"multline": False,
|
||||||
|
"value": "f-string",
|
||||||
|
},
|
||||||
|
"validate_template": {
|
||||||
|
"type": "bool",
|
||||||
|
"required": False,
|
||||||
|
"placeholder": "",
|
||||||
|
"list": False,
|
||||||
|
"show": False,
|
||||||
|
"multline": False,
|
||||||
|
"value": True,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"description": "Schema to represent a prompt for an LLM.",
|
||||||
|
"base_classes": ["BasePromptTemplate"],
|
||||||
|
}
|
||||||
|
|
||||||
|
return prompt
|
||||||
|
|
|
||||||
|
|
@ -7,6 +7,7 @@ from langchain import llms
|
||||||
from langchain.chains.conversation import memory as memories
|
from langchain.chains.conversation import memory as memories
|
||||||
from langchain.agents.load_tools import get_all_tool_names
|
from langchain.agents.load_tools import get_all_tool_names
|
||||||
from langflow.backend import util
|
from langflow.backend import util
|
||||||
|
from langflow.backend import customs
|
||||||
|
|
||||||
|
|
||||||
# build router
|
# build router
|
||||||
|
|
@ -52,10 +53,12 @@ def list_agents():
|
||||||
@router.get("/prompts")
|
@router.get("/prompts")
|
||||||
def list_prompts():
|
def list_prompts():
|
||||||
"""List all prompt types"""
|
"""List all prompt types"""
|
||||||
return [
|
custom_prompts = customs.get_custom_prompts()
|
||||||
|
library_prompts = [
|
||||||
prompt.__annotations__["return"].__name__
|
prompt.__annotations__["return"].__name__
|
||||||
for prompt in prompts.loading.type_to_loader_dict.values()
|
for prompt in prompts.loading.type_to_loader_dict.values()
|
||||||
]
|
]
|
||||||
|
return library_prompts + list(custom_prompts.keys())
|
||||||
|
|
||||||
|
|
||||||
@router.get("/llms")
|
@router.get("/llms")
|
||||||
|
|
|
||||||
|
|
@ -59,12 +59,16 @@ def build_json(root, nodes, edges):
|
||||||
# if module_type == "Tool":
|
# if module_type == "Tool":
|
||||||
# pass
|
# pass
|
||||||
if module_type in ["str", "bool", "int", "float", "Any"]:
|
if module_type in ["str", "bool", "int", "float", "Any"]:
|
||||||
|
# print(key)
|
||||||
|
# try:
|
||||||
value = value["value"]
|
value = value["value"]
|
||||||
|
# except:
|
||||||
|
# pass
|
||||||
elif "dict" in module_type:
|
elif "dict" in module_type:
|
||||||
value = {}
|
value = {}
|
||||||
else:
|
else:
|
||||||
# if value['list']:
|
# if value['list']:
|
||||||
print(key)
|
# print(key)
|
||||||
children = []
|
children = []
|
||||||
for c in local_nodes:
|
for c in local_nodes:
|
||||||
module_types = [c["data"]["type"]]
|
module_types = [c["data"]["type"]]
|
||||||
|
|
|
||||||
|
|
@ -10,6 +10,7 @@ from langchain.agents.load_tools import (
|
||||||
from langchain.chains.conversation import memory as memories
|
from langchain.chains.conversation import memory as memories
|
||||||
|
|
||||||
from langflow.backend import util
|
from langflow.backend import util
|
||||||
|
from langflow.backend import customs
|
||||||
|
|
||||||
# build router
|
# build router
|
||||||
router = APIRouter(
|
router = APIRouter(
|
||||||
|
|
@ -42,6 +43,8 @@ def get_agent(name: str):
|
||||||
def get_prompt(name: str):
|
def get_prompt(name: str):
|
||||||
"""Get the signature of a prompt."""
|
"""Get the signature of a prompt."""
|
||||||
try:
|
try:
|
||||||
|
if name in customs.get_custom_prompts().keys():
|
||||||
|
return customs.get_custom_prompts()[name]
|
||||||
return util.build_template_from_function(
|
return util.build_template_from_function(
|
||||||
name, prompts.loading.type_to_loader_dict
|
name, prompts.loading.type_to_loader_dict
|
||||||
)
|
)
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue