[Docs] use END instead of set_finish_point (#903)

This commit is contained in:
William FH
2024-07-01 21:56:10 -07:00
committed by GitHub
parent 727e63c01e
commit 320a87e1b9
31 changed files with 4189 additions and 389 deletions
+120 -9
View File
@@ -16,7 +16,10 @@
"id": "32a0e7f4",
"metadata": {},
"outputs": [],
"source": ["%%capture --no-stderr\n%pip install -U langgraph"]
"source": [
"%%capture --no-stderr\n",
"%pip install -U langgraph"
]
},
{
"cell_type": "markdown",
@@ -34,7 +37,73 @@
"id": "6d604311",
"metadata": {},
"outputs": [],
"source": ["import random\nfrom typing import Annotated, Literal\n\nfrom typing_extensions import TypedDict\n\nfrom langgraph.graph import StateGraph, START\nfrom langgraph.graph.message import add_messages\n\n\nclass State(TypedDict):\n messages: Annotated[list, add_messages]\n\n\nclass MyNode:\n def __init__(self, name: str):\n self.name = name\n\n def __call__(self, state: State):\n return {\"messages\": [(\"assistant\", f\"Called node {self.name}\")]}\n\n\ndef route(state) -> Literal[\"entry_node\", \"__end__\"]:\n if len(state[\"messages\"]) > 10:\n return \"__end__\"\n return \"entry_node\"\n\n\ndef add_fractal_nodes(builder, current_node, level, max_level):\n if level > max_level:\n return\n\n # Number of nodes to create at this level\n num_nodes = random.randint(1, 3) # Adjust randomness as needed\n for i in range(num_nodes):\n nm = [\"A\", \"B\", \"C\"][i]\n node_name = f\"node_{current_node}_{nm}\"\n builder.add_node(node_name, MyNode(node_name))\n builder.add_edge(current_node, node_name)\n\n # Recursively add more nodes\n r = random.random()\n if r > 0.2 and level + 1 < max_level:\n add_fractal_nodes(builder, node_name, level + 1, max_level)\n elif r > 0.05:\n builder.add_conditional_edges(node_name, route, node_name)\n else:\n # End\n builder.add_edge(node_name, \"__end__\")\n\n\ndef build_fractal_graph(max_level: int):\n builder = StateGraph(State)\n entry_point = \"entry_node\"\n builder.add_node(entry_point, MyNode(entry_point))\n builder.add_edge(START, entry_point)\n\n add_fractal_nodes(builder, entry_point, 1, max_level)\n\n # Optional: set a finish point if required\n builder.set_finish_point(entry_point) # or any specific node\n\n return builder.compile()\n\n\napp = build_fractal_graph(3)"]
"source": [
"import random\n",
"from typing import Annotated, Literal\n",
"\n",
"from typing_extensions import TypedDict\n",
"\n",
"from langgraph.graph import StateGraph, START, END\n",
"from langgraph.graph.message import add_messages\n",
"\n",
"\n",
"class State(TypedDict):\n",
" messages: Annotated[list, add_messages]\n",
"\n",
"\n",
"class MyNode:\n",
" def __init__(self, name: str):\n",
" self.name = name\n",
"\n",
" def __call__(self, state: State):\n",
" return {\"messages\": [(\"assistant\", f\"Called node {self.name}\")]}\n",
"\n",
"\n",
"def route(state) -> Literal[\"entry_node\", \"__end__\"]:\n",
" if len(state[\"messages\"]) > 10:\n",
" return \"__end__\"\n",
" return \"entry_node\"\n",
"\n",
"\n",
"def add_fractal_nodes(builder, current_node, level, max_level):\n",
" if level > max_level:\n",
" return\n",
"\n",
" # Number of nodes to create at this level\n",
" num_nodes = random.randint(1, 3) # Adjust randomness as needed\n",
" for i in range(num_nodes):\n",
" nm = [\"A\", \"B\", \"C\"][i]\n",
" node_name = f\"node_{current_node}_{nm}\"\n",
" builder.add_node(node_name, MyNode(node_name))\n",
" builder.add_edge(current_node, node_name)\n",
"\n",
" # Recursively add more nodes\n",
" r = random.random()\n",
" if r > 0.2 and level + 1 < max_level:\n",
" add_fractal_nodes(builder, node_name, level + 1, max_level)\n",
" elif r > 0.05:\n",
" builder.add_conditional_edges(node_name, route, node_name)\n",
" else:\n",
" # End\n",
" builder.add_edge(node_name, \"__end__\")\n",
"\n",
"\n",
"def build_fractal_graph(max_level: int):\n",
" builder = StateGraph(State)\n",
" entry_point = \"entry_node\"\n",
" builder.add_node(entry_point, MyNode(entry_point))\n",
" builder.add_edge(START, entry_point)\n",
"\n",
" add_fractal_nodes(builder, entry_point, 1, max_level)\n",
"\n",
" # Optional: set a finish point if required\n",
" builder.add_edge(entry_point, END) # or any specific node\n",
"\n",
" return builder.compile()\n",
"\n",
"\n",
"app = build_fractal_graph(3)"
]
},
{
"cell_type": "markdown",
@@ -96,7 +165,9 @@
]
}
],
"source": ["app.get_graph().print_ascii()"]
"source": [
"app.get_graph().print_ascii()"
]
},
{
"cell_type": "markdown",
@@ -155,7 +226,9 @@
]
}
],
"source": ["print(app.get_graph().draw_mermaid())"]
"source": [
"print(app.get_graph().draw_mermaid())"
]
},
{
"cell_type": "markdown",
@@ -193,7 +266,18 @@
"output_type": "display_data"
}
],
"source": ["from IPython.display import Image, display\nfrom langchain_core.runnables.graph import CurveStyle, MermaidDrawMethod, NodeColors\n\ndisplay(\n Image(\n app.get_graph().draw_mermaid_png(\n draw_method=MermaidDrawMethod.API,\n )\n )\n)"]
"source": [
"from IPython.display import Image, display\n",
"from langchain_core.runnables.graph import CurveStyle, MermaidDrawMethod, NodeColors\n",
"\n",
"display(\n",
" Image(\n",
" app.get_graph().draw_mermaid_png(\n",
" draw_method=MermaidDrawMethod.API,\n",
" )\n",
" )\n",
")"
]
},
{
"cell_type": "markdown",
@@ -219,7 +303,11 @@
}
},
"outputs": [],
"source": ["%%capture --no-stderr\n%pip install --quiet pyppeteer\n%pip install --quiet nest_asyncio"]
"source": [
"%%capture --no-stderr\n",
"%pip install --quiet pyppeteer\n",
"%pip install --quiet nest_asyncio"
]
},
{
"cell_type": "code",
@@ -243,7 +331,25 @@
"output_type": "display_data"
}
],
"source": ["import nest_asyncio\n\nnest_asyncio.apply() # Required for Jupyter Notebook to run async functions\n\ndisplay(\n Image(\n app.get_graph().draw_mermaid_png(\n curve_style=CurveStyle.LINEAR,\n node_colors=NodeColors(start=\"#ffdfba\", end=\"#baffc9\", other=\"#fad7de\"),\n wrap_label_n_words=9,\n output_file_path=None,\n draw_method=MermaidDrawMethod.PYPPETEER,\n background_color=\"white\",\n padding=10,\n )\n )\n)"]
"source": [
"import nest_asyncio\n",
"\n",
"nest_asyncio.apply() # Required for Jupyter Notebook to run async functions\n",
"\n",
"display(\n",
" Image(\n",
" app.get_graph().draw_mermaid_png(\n",
" curve_style=CurveStyle.LINEAR,\n",
" node_colors=NodeColors(start=\"#ffdfba\", end=\"#baffc9\", other=\"#fad7de\"),\n",
" wrap_label_n_words=9,\n",
" output_file_path=None,\n",
" draw_method=MermaidDrawMethod.PYPPETEER,\n",
" background_color=\"white\",\n",
" padding=10,\n",
" )\n",
" )\n",
")"
]
},
{
"cell_type": "markdown",
@@ -269,7 +375,10 @@
}
},
"outputs": [],
"source": ["%%capture --no-stderr\n%pip install pygraphviz"]
"source": [
"%%capture --no-stderr\n",
"%pip install pygraphviz"
]
},
{
"cell_type": "code",
@@ -293,7 +402,9 @@
"output_type": "display_data"
}
],
"source": ["display(Image(app.get_graph().draw_png()))"]
"source": [
"display(Image(app.get_graph().draw_png()))"
]
}
],
"metadata": {