[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
+11 -132
View File
@@ -32,10 +32,7 @@
"id": "8b323f43-328b-4b4b-88b0-6c84dc0a1d60",
"metadata": {},
"outputs": [],
"source": [
"%pip install -U --quiet langgraph langchain-fireworks\n",
"%pip install -U --quiet tavily-python"
]
"source": ["%pip install -U --quiet langgraph langchain-fireworks\n%pip install -U --quiet tavily-python"]
},
{
"cell_type": "code",
@@ -43,24 +40,7 @@
"id": "3368f330-cad6-4d35-a291-68fbf4389d98",
"metadata": {},
"outputs": [],
"source": [
"import getpass\n",
"import os\n",
"\n",
"\n",
"def _set_if_undefined(var: str) -> None:\n",
" if os.environ.get(var):\n",
" return\n",
" os.environ[var] = getpass.getpass(var)\n",
"\n",
"\n",
"# Optional: Configure tracing to visualize and debug the agent\n",
"_set_if_undefined(\"LANGCHAIN_API_KEY\")\n",
"os.environ[\"LANGCHAIN_TRACING_V2\"] = \"true\"\n",
"os.environ[\"LANGCHAIN_PROJECT\"] = \"Reflection\"\n",
"\n",
"_set_if_undefined(\"FIREWORKS_API_KEY\")"
]
"source": ["import getpass\nimport os\n\n\ndef _set_if_undefined(var: str) -> None:\n if os.environ.get(var):\n return\n os.environ[var] = getpass.getpass(var)\n\n\n# Optional: Configure tracing to visualize and debug the agent\n_set_if_undefined(\"LANGCHAIN_API_KEY\")\nos.environ[\"LANGCHAIN_TRACING_V2\"] = \"true\"\nos.environ[\"LANGCHAIN_PROJECT\"] = \"Reflection\"\n\n_set_if_undefined(\"FIREWORKS_API_KEY\")"]
},
{
"cell_type": "markdown",
@@ -78,28 +58,7 @@
"id": "cc10028f-9cef-4936-9419-cbdf06d24f1e",
"metadata": {},
"outputs": [],
"source": [
"from langchain_core.messages import AIMessage, BaseMessage, HumanMessage\n",
"from langchain_core.prompts import ChatPromptTemplate, MessagesPlaceholder\n",
"from langchain_fireworks import ChatFireworks\n",
"\n",
"prompt = ChatPromptTemplate.from_messages(\n",
" [\n",
" (\n",
" \"system\",\n",
" \"You are an essay assistant tasked with writing excellent 5-paragraph essays.\"\n",
" \" Generate the best essay possible for the user's request.\"\n",
" \" If the user provides critique, respond with a revised version of your previous attempts.\",\n",
" ),\n",
" MessagesPlaceholder(variable_name=\"messages\"),\n",
" ]\n",
")\n",
"llm = ChatFireworks(\n",
" model=\"accounts/fireworks/models/mixtral-8x7b-instruct\",\n",
" model_kwargs={\"max_tokens\": 32768},\n",
")\n",
"generate = prompt | llm"
]
"source": ["from langchain_core.messages import AIMessage, BaseMessage, HumanMessage\nfrom langchain_core.prompts import ChatPromptTemplate, MessagesPlaceholder\nfrom langchain_fireworks import ChatFireworks\n\nprompt = ChatPromptTemplate.from_messages(\n [\n (\n \"system\",\n \"You are an essay assistant tasked with writing excellent 5-paragraph essays.\"\n \" Generate the best essay possible for the user's request.\"\n \" If the user provides critique, respond with a revised version of your previous attempts.\",\n ),\n MessagesPlaceholder(variable_name=\"messages\"),\n ]\n)\nllm = ChatFireworks(\n model=\"accounts/fireworks/models/mixtral-8x7b-instruct\",\n model_kwargs={\"max_tokens\": 32768},\n)\ngenerate = prompt | llm"]
},
{
"cell_type": "code",
@@ -127,15 +86,7 @@
]
}
],
"source": [
"essay = \"\"\n",
"request = HumanMessage(\n",
" content=\"Write an essay on why the little prince is relevant in modern childhood\"\n",
")\n",
"for chunk in generate.stream({\"messages\": [request]}):\n",
" print(chunk.content, end=\"\")\n",
" essay += chunk.content"
]
"source": ["essay = \"\"\nrequest = HumanMessage(\n content=\"Write an essay on why the little prince is relevant in modern childhood\"\n)\nfor chunk in generate.stream({\"messages\": [request]}):\n print(chunk.content, end=\"\")\n essay += chunk.content"]
},
{
"cell_type": "markdown",
@@ -151,19 +102,7 @@
"id": "a705be92-88c0-4f4f-b4c2-cdcd9af8cb2c",
"metadata": {},
"outputs": [],
"source": [
"reflection_prompt = ChatPromptTemplate.from_messages(\n",
" [\n",
" (\n",
" \"system\",\n",
" \"You are a teacher grading an essay submission. Generate critique and recommendations for the user's submission.\"\n",
" \" Provide detailed recommendations, including requests for length, depth, style, etc.\",\n",
" ),\n",
" MessagesPlaceholder(variable_name=\"messages\"),\n",
" ]\n",
")\n",
"reflect = reflection_prompt | llm"
]
"source": ["reflection_prompt = ChatPromptTemplate.from_messages(\n [\n (\n \"system\",\n \"You are a teacher grading an essay submission. Generate critique and recommendations for the user's submission.\"\n \" Provide detailed recommendations, including requests for length, depth, style, etc.\",\n ),\n MessagesPlaceholder(variable_name=\"messages\"),\n ]\n)\nreflect = reflection_prompt | llm"]
},
{
"cell_type": "code",
@@ -193,12 +132,7 @@
]
}
],
"source": [
"reflection = \"\"\n",
"for chunk in reflect.stream({\"messages\": [request, HumanMessage(content=essay)]}):\n",
" print(chunk.content, end=\"\")\n",
" reflection += chunk.content"
]
"source": ["reflection = \"\"\nfor chunk in reflect.stream({\"messages\": [request, HumanMessage(content=essay)]}):\n print(chunk.content, end=\"\")\n reflection += chunk.content"]
},
{
"cell_type": "markdown",
@@ -236,12 +170,7 @@
]
}
],
"source": [
"for chunk in generate.stream(\n",
" {\"messages\": [request, AIMessage(content=essay), HumanMessage(content=reflection)]}\n",
"):\n",
" print(chunk.content, end=\"\")"
]
"source": ["for chunk in generate.stream(\n {\"messages\": [request, AIMessage(content=essay), HumanMessage(content=reflection)]}\n):\n print(chunk.content, end=\"\")"]
},
{
"cell_type": "markdown",
@@ -259,45 +188,7 @@
"id": "9e9a9d7c-5d2e-4194-b745-4511ec20db76",
"metadata": {},
"outputs": [],
"source": [
"from typing import List, Sequence\n",
"\n",
"from langgraph.graph import END, MessageGraph\n",
"\n",
"\n",
"async def generation_node(state: Sequence[BaseMessage]):\n",
" return await generate.ainvoke({\"messages\": state})\n",
"\n",
"\n",
"async def reflection_node(messages: Sequence[BaseMessage]) -> List[BaseMessage]:\n",
" # Other messages we need to adjust\n",
" cls_map = {\"ai\": HumanMessage, \"human\": AIMessage}\n",
" # First message is the original user request. We hold it the same for all nodes\n",
" translated = [messages[0]] + [\n",
" cls_map[msg.type](content=msg.content) for msg in messages[1:]\n",
" ]\n",
" res = await reflect.ainvoke({\"messages\": translated})\n",
" # We treat the output of this as human feedback for the generator\n",
" return HumanMessage(content=res.content)\n",
"\n",
"\n",
"builder = MessageGraph()\n",
"builder.add_node(\"generate\", generation_node)\n",
"builder.add_node(\"reflect\", reflection_node)\n",
"builder.set_entry_point(\"generate\")\n",
"\n",
"\n",
"def should_continue(state: List[BaseMessage]):\n",
" if len(state) > 6:\n",
" # End after 3 iterations\n",
" return END\n",
" return \"reflect\"\n",
"\n",
"\n",
"builder.add_conditional_edges(\"generate\", should_continue)\n",
"builder.add_edge(\"reflect\", \"generate\")\n",
"graph = builder.compile()"
]
"source": ["from typing import List, Sequence\n\nfrom langgraph.graph import END, MessageGraph, START\n\n\nasync def generation_node(state: Sequence[BaseMessage]):\n return await generate.ainvoke({\"messages\": state})\n\n\nasync def reflection_node(messages: Sequence[BaseMessage]) -> List[BaseMessage]:\n # Other messages we need to adjust\n cls_map = {\"ai\": HumanMessage, \"human\": AIMessage}\n # First message is the original user request. We hold it the same for all nodes\n translated = [messages[0]] + [\n cls_map[msg.type](content=msg.content) for msg in messages[1:]\n ]\n res = await reflect.ainvoke({\"messages\": translated})\n # We treat the output of this as human feedback for the generator\n return HumanMessage(content=res.content)\n\n\nbuilder = MessageGraph()\nbuilder.add_node(\"generate\", generation_node)\nbuilder.add_node(\"reflect\", reflection_node)\nbuilder.add_edge(START, \"generate\")\n\n\ndef should_continue(state: List[BaseMessage]):\n if len(state) > 6:\n # End after 3 iterations\n return END\n return \"reflect\"\n\n\nbuilder.add_conditional_edges(\"generate\", should_continue)\nbuilder.add_edge(\"reflect\", \"generate\")\ngraph = builder.compile()"]
},
{
"cell_type": "code",
@@ -328,17 +219,7 @@
]
}
],
"source": [
"async for event in graph.astream(\n",
" [\n",
" HumanMessage(\n",
" content=\"Generate an essay on the topicality of The Little Prince and its message in modern life\"\n",
" )\n",
" ],\n",
"):\n",
" print(event)\n",
" print(\"---\")"
]
"source": ["async for event in graph.astream(\n [\n HumanMessage(\n content=\"Generate an essay on the topicality of The Little Prince and its message in modern life\"\n )\n ],\n):\n print(event)\n print(\"---\")"]
},
{
"cell_type": "code",
@@ -490,9 +371,7 @@
]
}
],
"source": [
"ChatPromptTemplate.from_messages(event[END]).pretty_print()"
]
"source": ["ChatPromptTemplate.from_messages(event[END]).pretty_print()"]
},
{
"cell_type": "markdown",
@@ -510,7 +389,7 @@
"id": "7c0e3efd-7f54-410e-bd31-36185a46b9a8",
"metadata": {},
"outputs": [],
"source": []
"source": [""]
}
],
"metadata": {