[Docs] Update notebooks to use START (#902)

This commit is contained in:
William FH
2024-07-01 21:36:34 -07:00
committed by GitHub
parent 267f5e5234
commit 727e63c01e
67 changed files with 1059 additions and 20258 deletions
+9 -120
View File
@@ -16,10 +16,7 @@
"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",
@@ -37,73 +34,7 @@
"id": "6d604311",
"metadata": {},
"outputs": [],
"source": [
"import random\n",
"from typing import Annotated, Literal\n",
"\n",
"from typing_extensions import TypedDict\n",
"\n",
"from langgraph.graph import StateGraph\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.set_entry_point(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",
"\n",
"app = build_fractal_graph(3)"
]
"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)"]
},
{
"cell_type": "markdown",
@@ -165,9 +96,7 @@
]
}
],
"source": [
"app.get_graph().print_ascii()"
]
"source": ["app.get_graph().print_ascii()"]
},
{
"cell_type": "markdown",
@@ -226,9 +155,7 @@
]
}
],
"source": [
"print(app.get_graph().draw_mermaid())"
]
"source": ["print(app.get_graph().draw_mermaid())"]
},
{
"cell_type": "markdown",
@@ -266,18 +193,7 @@
"output_type": "display_data"
}
],
"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",
")"
]
"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)"]
},
{
"cell_type": "markdown",
@@ -303,11 +219,7 @@
}
},
"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",
@@ -331,25 +243,7 @@
"output_type": "display_data"
}
],
"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",
")"
]
"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)"]
},
{
"cell_type": "markdown",
@@ -375,10 +269,7 @@
}
},
"outputs": [],
"source": [
"%%capture --no-stderr\n",
"%pip install pygraphviz"
]
"source": ["%%capture --no-stderr\n%pip install pygraphviz"]
},
{
"cell_type": "code",
@@ -402,9 +293,7 @@
"output_type": "display_data"
}
],
"source": [
"display(Image(app.get_graph().draw_png()))"
]
"source": ["display(Image(app.get_graph().draw_png()))"]
}
],
"metadata": {