[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
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
+186 -16
View File
@@ -56,7 +56,9 @@
"id": "4a660963-bd3d-4c87-b2e4-b6e432055211",
"metadata": {},
"outputs": [],
"source": ["%%capture --no-stderr\n%pip install -U langchain_community tiktoken langchainhub scikit-learn langchain langgraph tavily-python nomic[local] langchain-nomic langchain_openai"]
"source": [
"%%capture --no-stderr\n%pip install -U langchain_community tiktoken langchainhub scikit-learn langchain langgraph tavily-python nomic[local] langchain-nomic langchain_openai"
]
},
{
"cell_type": "code",
@@ -64,7 +66,12 @@
"id": "68316ba0-854b-41e1-9af5-1f9e965946e3",
"metadata": {},
"outputs": [],
"source": ["# Search\nimport os\nos.environ[\"TAVILY_API_KEY\"] = \"xxx\""]
"source": [
"# Search\n",
"import os\n",
"\n",
"os.environ[\"TAVILY_API_KEY\"] = \"xxx\""
]
},
{
"cell_type": "code",
@@ -72,7 +79,9 @@
"id": "0be68860-dded-481e-9fc7-a5042bf92c04",
"metadata": {},
"outputs": [],
"source": ["# Embedding (optional)\nos.environ[\"OPENAI_API_KEY\"] = \"xxx\""]
"source": [
"# Embedding (optional)\nos.environ[\"OPENAI_API_KEY\"] = \"xxx\""
]
},
{
"cell_type": "code",
@@ -80,7 +89,9 @@
"id": "7248ab88-2b97-41eb-8dbb-4ea65525ed9a",
"metadata": {},
"outputs": [],
"source": ["# Tracing and testing (optional)\nos.environ[\"LANGCHAIN_API_KEY\"] = \"xxx\"\nos.environ[\"LANGCHAIN_TRACING_V2\"] = \"true\"\nos.environ[\"LANGCHAIN_ENDPOINT\"] = \"https://api.smith.langchain.com\"\nos.environ[\"LANGCHAIN_PROJECT\"] = \"corrective-rag-agent-testing\""]
"source": [
"# Tracing and testing (optional)\nos.environ[\"LANGCHAIN_API_KEY\"] = \"xxx\"\nos.environ[\"LANGCHAIN_TRACING_V2\"] = \"true\"\nos.environ[\"LANGCHAIN_ENDPOINT\"] = \"https://api.smith.langchain.com\"\nos.environ[\"LANGCHAIN_PROJECT\"] = \"corrective-rag-agent-testing\""
]
},
{
"cell_type": "markdown",
@@ -98,7 +109,9 @@
"id": "2f4db331-c4d0-4c7c-a9a5-0bebc8a89c6c",
"metadata": {},
"outputs": [],
"source": ["local_llm = \"llama3\"\nmodel_tested = \"llama3-8b\"\nmetadata = f\"CRAG, {model_tested}\""]
"source": [
"local_llm = \"llama3\"\nmodel_tested = \"llama3-8b\"\nmetadata = f\"CRAG, {model_tested}\""
]
},
{
"cell_type": "markdown",
@@ -124,7 +137,48 @@
]
}
],
"source": ["from langchain.text_splitter import RecursiveCharacterTextSplitter\nfrom langchain_community.document_loaders import WebBaseLoader\nfrom langchain_community.vectorstores import SKLearnVectorStore\nfrom langchain_nomic.embeddings import NomicEmbeddings # local\nfrom langchain_openai import OpenAIEmbeddings # api\n\n# List of URLs to load documents from\nurls = [\n \"https://lilianweng.github.io/posts/2023-06-23-agent/\",\n \"https://lilianweng.github.io/posts/2023-03-15-prompt-engineering/\",\n \"https://lilianweng.github.io/posts/2023-10-25-adv-attack-llm/\",\n]\n\n# Load documents from the URLs\ndocs = [WebBaseLoader(url).load() for url in urls]\ndocs_list = [item for sublist in docs for item in sublist]\n\n# Initialize a text splitter with specified chunk size and overlap\ntext_splitter = RecursiveCharacterTextSplitter.from_tiktoken_encoder(\n chunk_size=250, chunk_overlap=0\n)\n\n# Split the documents into chunks\ndoc_splits = text_splitter.split_documents(docs_list)\n\n# Embedding\n'''\nembedding=NomicEmbeddings(\n model=\"nomic-embed-text-v1.5\",\n inference_mode=\"local\",\n)\n'''\nembedding = OpenAIEmbeddings()\n\n# Add the document chunks to the \"vector store\"\nvectorstore = SKLearnVectorStore.from_documents(\n documents=doc_splits,\n embedding=embedding,\n)\nretriever = vectorstore.as_retriever(k=4)"]
"source": [
"from langchain.text_splitter import RecursiveCharacterTextSplitter\n",
"from langchain_community.document_loaders import WebBaseLoader\n",
"from langchain_community.vectorstores import SKLearnVectorStore\n",
"from langchain_nomic.embeddings import NomicEmbeddings # local\n",
"from langchain_openai import OpenAIEmbeddings # api\n",
"\n",
"# List of URLs to load documents from\n",
"urls = [\n",
" \"https://lilianweng.github.io/posts/2023-06-23-agent/\",\n",
" \"https://lilianweng.github.io/posts/2023-03-15-prompt-engineering/\",\n",
" \"https://lilianweng.github.io/posts/2023-10-25-adv-attack-llm/\",\n",
"]\n",
"\n",
"# Load documents from the URLs\n",
"docs = [WebBaseLoader(url).load() for url in urls]\n",
"docs_list = [item for sublist in docs for item in sublist]\n",
"\n",
"# Initialize a text splitter with specified chunk size and overlap\n",
"text_splitter = RecursiveCharacterTextSplitter.from_tiktoken_encoder(\n",
" chunk_size=250, chunk_overlap=0\n",
")\n",
"\n",
"# Split the documents into chunks\n",
"doc_splits = text_splitter.split_documents(docs_list)\n",
"\n",
"# Embedding\n",
"\"\"\"\n",
"embedding=NomicEmbeddings(\n",
" model=\"nomic-embed-text-v1.5\",\n",
" inference_mode=\"local\",\n",
")\n",
"\"\"\"\n",
"embedding = OpenAIEmbeddings()\n",
"\n",
"# Add the document chunks to the \"vector store\"\n",
"vectorstore = SKLearnVectorStore.from_documents(\n",
" documents=doc_splits,\n",
" embedding=embedding,\n",
")\n",
"retriever = vectorstore.as_retriever(k=4)"
]
},
{
"attachments": {},
@@ -149,7 +203,9 @@
]
}
],
"source": ["### Retrieval Grader\n\nfrom langchain.prompts import PromptTemplate\nfrom langchain_community.chat_models import ChatOllama\nfrom langchain_core.output_parsers import JsonOutputParser\nfrom langchain_mistralai.chat_models import ChatMistralAI\n\n# LLM\nllm = ChatOllama(model=local_llm, format=\"json\", temperature=0)\n\n# Prompt\nprompt = PromptTemplate(\n template=\"\"\"You are a teacher grading a quiz. You will be given: \n 1/ a QUESTION\n 2/ A FACT provided by the student\n \n You are grading RELEVANCE RECALL:\n A score of 1 means that ANY of the statements in the FACT are relevant to the QUESTION. \n A score of 0 means that NONE of the statements in the FACT are relevant to the QUESTION. \n 1 is the highest (best) score. 0 is the lowest score you can give. \n \n Explain your reasoning in a step-by-step manner. Ensure your reasoning and conclusion are correct. \n \n Avoid simply stating the correct answer at the outset.\n \n Question: {question} \\n\n Fact: \\n\\n {documents} \\n\\n\n \n Give a binary score 'yes' or 'no' score to indicate whether the document is relevant to the question. \\n\n Provide the binary score as a JSON with a single key 'score' and no premable or explanation.\n \"\"\",\n input_variables=[\"question\", \"documents\"],\n)\n\nretrieval_grader = prompt | llm | JsonOutputParser()\nquestion = \"agent memory\"\ndocs = retriever.invoke(question)\ndoc_txt = docs[1].page_content\nprint(retrieval_grader.invoke({\"question\": question, \"documents\": doc_txt}))"]
"source": [
"### Retrieval Grader\n\nfrom langchain.prompts import PromptTemplate\nfrom langchain_community.chat_models import ChatOllama\nfrom langchain_core.output_parsers import JsonOutputParser\nfrom langchain_mistralai.chat_models import ChatMistralAI\n\n# LLM\nllm = ChatOllama(model=local_llm, format=\"json\", temperature=0)\n\n# Prompt\nprompt = PromptTemplate(\n template=\"\"\"You are a teacher grading a quiz. You will be given: \n 1/ a QUESTION\n 2/ A FACT provided by the student\n \n You are grading RELEVANCE RECALL:\n A score of 1 means that ANY of the statements in the FACT are relevant to the QUESTION. \n A score of 0 means that NONE of the statements in the FACT are relevant to the QUESTION. \n 1 is the highest (best) score. 0 is the lowest score you can give. \n \n Explain your reasoning in a step-by-step manner. Ensure your reasoning and conclusion are correct. \n \n Avoid simply stating the correct answer at the outset.\n \n Question: {question} \\n\n Fact: \\n\\n {documents} \\n\\n\n \n Give a binary score 'yes' or 'no' score to indicate whether the document is relevant to the question. \\n\n Provide the binary score as a JSON with a single key 'score' and no premable or explanation.\n \"\"\",\n input_variables=[\"question\", \"documents\"],\n)\n\nretrieval_grader = prompt | llm | JsonOutputParser()\nquestion = \"agent memory\"\ndocs = retriever.invoke(question)\ndoc_txt = docs[1].page_content\nprint(retrieval_grader.invoke({\"question\": question, \"documents\": doc_txt}))"
]
},
{
"cell_type": "code",
@@ -165,7 +221,9 @@
]
}
],
"source": ["### Generate\n\nfrom langchain_core.output_parsers import StrOutputParser\n\n# Prompt\nprompt = PromptTemplate(\n template=\"\"\"You are an assistant for question-answering tasks. \n \n Use the following documents to answer the question. \n \n If you don't know the answer, just say that you don't know. \n \n Use three sentences maximum and keep the answer concise:\n Question: {question} \n Documents: {documents} \n Answer: \n \"\"\",\n input_variables=[\"question\", \"documents\"],\n)\n\n# LLM\nllm = ChatOllama(model=local_llm, temperature=0)\n\n# Chain\nrag_chain = prompt | llm | StrOutputParser()\n\n# Run\ngeneration = rag_chain.invoke({\"documents\": docs, \"question\": question})\nprint(generation)"]
"source": [
"### Generate\n\nfrom langchain_core.output_parsers import StrOutputParser\n\n# Prompt\nprompt = PromptTemplate(\n template=\"\"\"You are an assistant for question-answering tasks. \n \n Use the following documents to answer the question. \n \n If you don't know the answer, just say that you don't know. \n \n Use three sentences maximum and keep the answer concise:\n Question: {question} \n Documents: {documents} \n Answer: \n \"\"\",\n input_variables=[\"question\", \"documents\"],\n)\n\n# LLM\nllm = ChatOllama(model=local_llm, temperature=0)\n\n# Chain\nrag_chain = prompt | llm | StrOutputParser()\n\n# Run\ngeneration = rag_chain.invoke({\"documents\": docs, \"question\": question})\nprint(generation)"
]
},
{
"cell_type": "code",
@@ -173,7 +231,9 @@
"id": "b36a2f36-bc5f-408d-a5e8-3fa203c233f6",
"metadata": {},
"outputs": [],
"source": ["### Search\n\nfrom langchain_community.tools.tavily_search import TavilySearchResults\n\nweb_search_tool = TavilySearchResults(k=3)"]
"source": [
"### Search\n\nfrom langchain_community.tools.tavily_search import TavilySearchResults\n\nweb_search_tool = TavilySearchResults(k=3)"
]
},
{
"cell_type": "markdown",
@@ -202,7 +262,9 @@
"output_type": "display_data"
}
],
"source": ["from typing import List\nfrom typing_extensions import TypedDict\nfrom IPython.display import Image, display\nfrom langchain.schema import Document\nfrom langgraph.graph import START, END, StateGraph\n\n\nclass GraphState(TypedDict):\n \"\"\"\n Represents the state of our graph.\n\n Attributes:\n question: question\n generation: LLM generation\n search: whether to add search\n documents: list of documents\n \"\"\"\n\n question: str\n generation: str\n search: str\n documents: List[str]\n steps: List[str]\n\n\ndef retrieve(state):\n \"\"\"\n Retrieve documents\n\n Args:\n state (dict): The current graph state\n\n Returns:\n state (dict): New key added to state, documents, that contains retrieved documents\n \"\"\"\n question = state[\"question\"]\n documents = retriever.invoke(question)\n steps = state[\"steps\"]\n steps.append(\"retrieve_documents\")\n return {\"documents\": documents, \"question\": question, \"steps\": steps}\n\n\ndef generate(state):\n \"\"\"\n Generate answer\n\n Args:\n state (dict): The current graph state\n\n Returns:\n state (dict): New key added to state, generation, that contains LLM generation\n \"\"\"\n\n question = state[\"question\"]\n documents = state[\"documents\"]\n generation = rag_chain.invoke({\"documents\": documents, \"question\": question})\n steps = state[\"steps\"]\n steps.append(\"generate_answer\")\n return {\n \"documents\": documents,\n \"question\": question,\n \"generation\": generation,\n \"steps\": steps,\n }\n\n\ndef grade_documents(state):\n \"\"\"\n Determines whether the retrieved documents are relevant to the question.\n\n Args:\n state (dict): The current graph state\n\n Returns:\n state (dict): Updates documents key with only filtered relevant documents\n \"\"\"\n\n question = state[\"question\"]\n documents = state[\"documents\"]\n steps = state[\"steps\"]\n steps.append(\"grade_document_retrieval\")\n filtered_docs = []\n search = \"No\"\n for d in documents:\n score = retrieval_grader.invoke(\n {\"question\": question, \"documents\": d.page_content}\n )\n grade = score[\"score\"]\n if grade == \"yes\":\n filtered_docs.append(d)\n else:\n search = \"Yes\"\n continue\n return {\n \"documents\": filtered_docs,\n \"question\": question,\n \"search\": search,\n \"steps\": steps,\n }\n\n\ndef web_search(state):\n \"\"\"\n Web search based on the re-phrased question.\n\n Args:\n state (dict): The current graph state\n\n Returns:\n state (dict): Updates documents key with appended web results\n \"\"\"\n\n question = state[\"question\"]\n documents = state.get(\"documents\", [])\n steps = state[\"steps\"]\n steps.append(\"web_search\")\n web_results = web_search_tool.invoke({\"query\": question})\n documents.extend(\n [\n Document(page_content=d[\"content\"], metadata={\"url\": d[\"url\"]})\n for d in web_results\n ]\n )\n return {\"documents\": documents, \"question\": question, \"steps\": steps}\n\n\ndef decide_to_generate(state):\n \"\"\"\n Determines whether to generate an answer, or re-generate a question.\n\n Args:\n state (dict): The current graph state\n\n Returns:\n str: Binary decision for next node to call\n \"\"\"\n search = state[\"search\"]\n if search == \"Yes\":\n return \"search\"\n else:\n return \"generate\"\n\n\n# Graph\nworkflow = StateGraph(GraphState)\n\n# Define the nodes\nworkflow.add_node(\"retrieve\", retrieve) # retrieve\nworkflow.add_node(\"grade_documents\", grade_documents) # grade documents\nworkflow.add_node(\"generate\", generate) # generatae\nworkflow.add_node(\"web_search\", web_search) # web search\n\n# Build graph\nworkflow.add_edge(START, \"retrieve\")\nworkflow.add_edge(\"retrieve\", \"grade_documents\")\nworkflow.add_conditional_edges(\n \"grade_documents\",\n decide_to_generate,\n {\n \"search\": \"web_search\",\n \"generate\": \"generate\",\n },\n)\nworkflow.add_edge(\"web_search\", \"generate\")\nworkflow.add_edge(\"generate\", END)\n\ncustom_graph = workflow.compile()\n\ndisplay(Image(custom_graph.get_graph(xray=True).draw_mermaid_png()))"]
"source": [
"from typing import List\nfrom typing_extensions import TypedDict\nfrom IPython.display import Image, display\nfrom langchain.schema import Document\nfrom langgraph.graph import START, END, StateGraph\n\n\nclass GraphState(TypedDict):\n \"\"\"\n Represents the state of our graph.\n\n Attributes:\n question: question\n generation: LLM generation\n search: whether to add search\n documents: list of documents\n \"\"\"\n\n question: str\n generation: str\n search: str\n documents: List[str]\n steps: List[str]\n\n\ndef retrieve(state):\n \"\"\"\n Retrieve documents\n\n Args:\n state (dict): The current graph state\n\n Returns:\n state (dict): New key added to state, documents, that contains retrieved documents\n \"\"\"\n question = state[\"question\"]\n documents = retriever.invoke(question)\n steps = state[\"steps\"]\n steps.append(\"retrieve_documents\")\n return {\"documents\": documents, \"question\": question, \"steps\": steps}\n\n\ndef generate(state):\n \"\"\"\n Generate answer\n\n Args:\n state (dict): The current graph state\n\n Returns:\n state (dict): New key added to state, generation, that contains LLM generation\n \"\"\"\n\n question = state[\"question\"]\n documents = state[\"documents\"]\n generation = rag_chain.invoke({\"documents\": documents, \"question\": question})\n steps = state[\"steps\"]\n steps.append(\"generate_answer\")\n return {\n \"documents\": documents,\n \"question\": question,\n \"generation\": generation,\n \"steps\": steps,\n }\n\n\ndef grade_documents(state):\n \"\"\"\n Determines whether the retrieved documents are relevant to the question.\n\n Args:\n state (dict): The current graph state\n\n Returns:\n state (dict): Updates documents key with only filtered relevant documents\n \"\"\"\n\n question = state[\"question\"]\n documents = state[\"documents\"]\n steps = state[\"steps\"]\n steps.append(\"grade_document_retrieval\")\n filtered_docs = []\n search = \"No\"\n for d in documents:\n score = retrieval_grader.invoke(\n {\"question\": question, \"documents\": d.page_content}\n )\n grade = score[\"score\"]\n if grade == \"yes\":\n filtered_docs.append(d)\n else:\n search = \"Yes\"\n continue\n return {\n \"documents\": filtered_docs,\n \"question\": question,\n \"search\": search,\n \"steps\": steps,\n }\n\n\ndef web_search(state):\n \"\"\"\n Web search based on the re-phrased question.\n\n Args:\n state (dict): The current graph state\n\n Returns:\n state (dict): Updates documents key with appended web results\n \"\"\"\n\n question = state[\"question\"]\n documents = state.get(\"documents\", [])\n steps = state[\"steps\"]\n steps.append(\"web_search\")\n web_results = web_search_tool.invoke({\"query\": question})\n documents.extend(\n [\n Document(page_content=d[\"content\"], metadata={\"url\": d[\"url\"]})\n for d in web_results\n ]\n )\n return {\"documents\": documents, \"question\": question, \"steps\": steps}\n\n\ndef decide_to_generate(state):\n \"\"\"\n Determines whether to generate an answer, or re-generate a question.\n\n Args:\n state (dict): The current graph state\n\n Returns:\n str: Binary decision for next node to call\n \"\"\"\n search = state[\"search\"]\n if search == \"Yes\":\n return \"search\"\n else:\n return \"generate\"\n\n\n# Graph\nworkflow = StateGraph(GraphState)\n\n# Define the nodes\nworkflow.add_node(\"retrieve\", retrieve) # retrieve\nworkflow.add_node(\"grade_documents\", grade_documents) # grade documents\nworkflow.add_node(\"generate\", generate) # generatae\nworkflow.add_node(\"web_search\", web_search) # web search\n\n# Build graph\nworkflow.add_edge(START, \"retrieve\")\nworkflow.add_edge(\"retrieve\", \"grade_documents\")\nworkflow.add_conditional_edges(\n \"grade_documents\",\n decide_to_generate,\n {\n \"search\": \"web_search\",\n \"generate\": \"generate\",\n },\n)\nworkflow.add_edge(\"web_search\", \"generate\")\nworkflow.add_edge(\"generate\", END)\n\ncustom_graph = workflow.compile()\n\ndisplay(Image(custom_graph.get_graph(xray=True).draw_mermaid_png()))"
]
},
{
"cell_type": "code",
@@ -225,7 +287,22 @@
"output_type": "execute_result"
}
],
"source": ["import uuid\n\ndef predict_custom_agent_local_answer(example: dict):\n config = {\"configurable\": {\"thread_id\": str(uuid.uuid4())}}\n state_dict = custom_graph.invoke(\n {\"question\": example[\"input\"], \"steps\": []}, config\n )\n return {\"response\": state_dict[\"generation\"], \"steps\": state_dict[\"steps\"]}\n\n\nexample = {\"input\": \"What are the types of agent memory?\"}\nresponse = predict_custom_agent_local_answer(example)\nresponse"]
"source": [
"import uuid\n",
"\n",
"\n",
"def predict_custom_agent_local_answer(example: dict):\n",
" config = {\"configurable\": {\"thread_id\": str(uuid.uuid4())}}\n",
" state_dict = custom_graph.invoke(\n",
" {\"question\": example[\"input\"], \"steps\": []}, config\n",
" )\n",
" return {\"response\": state_dict[\"generation\"], \"steps\": state_dict[\"steps\"]}\n",
"\n",
"\n",
"example = {\"input\": \"What are the types of agent memory?\"}\n",
"response = predict_custom_agent_local_answer(example)\n",
"response"
]
},
{
"cell_type": "markdown",
@@ -261,7 +338,9 @@
"id": "b83706ac-724b-46b1-9f08-66e6c4fac742",
"metadata": {},
"outputs": [],
"source": ["from langsmith import Client\n\nclient = Client()\n\n# Create a dataset\nexamples = [\n (\n \"How does the ReAct agent use self-reflection? \",\n \"ReAct integrates reasoning and acting, performing actions - such tools like Wikipedia search API - and then observing / reasoning about the tool outputs.\",\n ),\n (\n \"What are the types of biases that can arise with few-shot prompting?\",\n \"The biases that can arise with few-shot prompting include (1) Majority label bias, (2) Recency bias, and (3) Common token bias.\",\n ),\n (\n \"What are five types of adversarial attacks?\",\n \"Five types of adversarial attacks are (1) Token manipulation, (2) Gradient based attack, (3) Jailbreak prompting, (4) Human red-teaming, (5) Model red-teaming.\",\n ),\n (\n \"Who did the Chicago Bears draft first in the 2024 NFL draft”?\",\n \"The Chicago Bears drafted Caleb Williams first in the 2024 NFL draft.\",\n ),\n (\"Who won the 2024 NBA finals?\", \"The Boston Celtics on the 2024 NBA finals\"),\n]\n\n# Save it\ndataset_name = \"Corrective RAG Agent Testing\"\nif not client.has_dataset(dataset_name=dataset_name):\n dataset = client.create_dataset(dataset_name=dataset_name)\n inputs, outputs = zip(\n *[({\"input\": text}, {\"output\": label}) for text, label in examples]\n )\n client.create_examples(inputs=inputs, outputs=outputs, dataset_id=dataset.id)"]
"source": [
"from langsmith import Client\n\nclient = Client()\n\n# Create a dataset\nexamples = [\n (\n \"How does the ReAct agent use self-reflection? \",\n \"ReAct integrates reasoning and acting, performing actions - such tools like Wikipedia search API - and then observing / reasoning about the tool outputs.\",\n ),\n (\n \"What are the types of biases that can arise with few-shot prompting?\",\n \"The biases that can arise with few-shot prompting include (1) Majority label bias, (2) Recency bias, and (3) Common token bias.\",\n ),\n (\n \"What are five types of adversarial attacks?\",\n \"Five types of adversarial attacks are (1) Token manipulation, (2) Gradient based attack, (3) Jailbreak prompting, (4) Human red-teaming, (5) Model red-teaming.\",\n ),\n (\n \"Who did the Chicago Bears draft first in the 2024 NFL draft”?\",\n \"The Chicago Bears drafted Caleb Williams first in the 2024 NFL draft.\",\n ),\n (\"Who won the 2024 NBA finals?\", \"The Boston Celtics on the 2024 NBA finals\"),\n]\n\n# Save it\ndataset_name = \"Corrective RAG Agent Testing\"\nif not client.has_dataset(dataset_name=dataset_name):\n dataset = client.create_dataset(dataset_name=dataset_name)\n inputs, outputs = zip(\n *[({\"input\": text}, {\"output\": label}) for text, label in examples]\n )\n client.create_examples(inputs=inputs, outputs=outputs, dataset_id=dataset.id)"
]
},
{
"cell_type": "markdown",
@@ -281,7 +360,39 @@
"id": "0a63776c-f9cd-46ce-b8cf-95c066dc5b06",
"metadata": {},
"outputs": [],
"source": ["from langchain import hub\nfrom langchain_openai import ChatOpenAI\n\n# Grade prompt\ngrade_prompt_answer_accuracy = hub.pull(\"langchain-ai/rag-answer-vs-reference\")\n\ndef answer_evaluator(run, example) -> dict:\n \"\"\"\n A simple evaluator for RAG answer accuracy\n \"\"\"\n\n # Get the question, the ground truth reference answer, RAG chain answer prediction\n input_question = example.inputs[\"input\"]\n reference = example.outputs[\"output\"]\n prediction = run.outputs[\"response\"]\n\n # Define an LLM grader\n llm = ChatOpenAI(model=\"gpt-4o\", temperature=0)\n answer_grader = grade_prompt_answer_accuracy | llm\n\n # Run evaluator\n score = answer_grader.invoke(\n {\n \"question\": input_question,\n \"correct_answer\": reference,\n \"student_answer\": prediction,\n }\n )\n score = score[\"Score\"]\n return {\"key\": \"answer_v_reference_score\", \"score\": score}"]
"source": [
"from langchain import hub\n",
"from langchain_openai import ChatOpenAI\n",
"\n",
"# Grade prompt\n",
"grade_prompt_answer_accuracy = hub.pull(\"langchain-ai/rag-answer-vs-reference\")\n",
"\n",
"\n",
"def answer_evaluator(run, example) -> dict:\n",
" \"\"\"\n",
" A simple evaluator for RAG answer accuracy\n",
" \"\"\"\n",
"\n",
" # Get the question, the ground truth reference answer, RAG chain answer prediction\n",
" input_question = example.inputs[\"input\"]\n",
" reference = example.outputs[\"output\"]\n",
" prediction = run.outputs[\"response\"]\n",
"\n",
" # Define an LLM grader\n",
" llm = ChatOpenAI(model=\"gpt-4o\", temperature=0)\n",
" answer_grader = grade_prompt_answer_accuracy | llm\n",
"\n",
" # Run evaluator\n",
" score = answer_grader.invoke(\n",
" {\n",
" \"question\": input_question,\n",
" \"correct_answer\": reference,\n",
" \"student_answer\": prediction,\n",
" }\n",
" )\n",
" score = score[\"Score\"]\n",
" return {\"key\": \"answer_v_reference_score\", \"score\": score}"
]
},
{
"cell_type": "markdown",
@@ -301,7 +412,51 @@
"id": "deb28175-27a1-4afc-9747-2983e87fc881",
"metadata": {},
"outputs": [],
"source": ["from langsmith.schemas import Example, Run\n\n# Reasoning traces that we expect the agents to take\nexpected_trajectory_1 = [\n \"retrieve_documents\",\n \"grade_document_retrieval\",\n \"web_search\",\n \"generate_answer\",\n]\nexpected_trajectory_2 = [\n \"retrieve_documents\",\n \"grade_document_retrieval\",\n \"generate_answer\",\n]\n\ndef check_trajectory_react(root_run: Run, example: Example) -> dict:\n \"\"\"\n Check if all expected tools are called in exact order and without any additional tool calls.\n \"\"\"\n messages = root_run.outputs[\"messages\"]\n tool_calls = find_tool_calls_react(messages)\n print(f\"Tool calls ReAct agent: {tool_calls}\")\n if tool_calls == expected_trajectory_1 or tool_calls == expected_trajectory_2:\n score = 1\n else:\n score = 0\n\n return {\"score\": int(score), \"key\": \"tool_calls_in_exact_order\"}\n\n\ndef check_trajectory_custom(root_run: Run, example: Example) -> dict:\n \"\"\"\n Check if all expected tools are called in exact order and without any additional tool calls.\n \"\"\"\n tool_calls = root_run.outputs[\"steps\"]\n print(f\"Tool calls custom agent: {tool_calls}\")\n if tool_calls == expected_trajectory_1 or tool_calls == expected_trajectory_2:\n score = 1\n else:\n score = 0\n\n return {\"score\": int(score), \"key\": \"tool_calls_in_exact_order\"}"]
"source": [
"from langsmith.schemas import Example, Run\n",
"\n",
"# Reasoning traces that we expect the agents to take\n",
"expected_trajectory_1 = [\n",
" \"retrieve_documents\",\n",
" \"grade_document_retrieval\",\n",
" \"web_search\",\n",
" \"generate_answer\",\n",
"]\n",
"expected_trajectory_2 = [\n",
" \"retrieve_documents\",\n",
" \"grade_document_retrieval\",\n",
" \"generate_answer\",\n",
"]\n",
"\n",
"\n",
"def check_trajectory_react(root_run: Run, example: Example) -> dict:\n",
" \"\"\"\n",
" Check if all expected tools are called in exact order and without any additional tool calls.\n",
" \"\"\"\n",
" messages = root_run.outputs[\"messages\"]\n",
" tool_calls = find_tool_calls_react(messages)\n",
" print(f\"Tool calls ReAct agent: {tool_calls}\")\n",
" if tool_calls == expected_trajectory_1 or tool_calls == expected_trajectory_2:\n",
" score = 1\n",
" else:\n",
" score = 0\n",
"\n",
" return {\"score\": int(score), \"key\": \"tool_calls_in_exact_order\"}\n",
"\n",
"\n",
"def check_trajectory_custom(root_run: Run, example: Example) -> dict:\n",
" \"\"\"\n",
" Check if all expected tools are called in exact order and without any additional tool calls.\n",
" \"\"\"\n",
" tool_calls = root_run.outputs[\"steps\"]\n",
" print(f\"Tool calls custom agent: {tool_calls}\")\n",
" if tool_calls == expected_trajectory_1 or tool_calls == expected_trajectory_2:\n",
" score = 1\n",
" else:\n",
" score = 0\n",
"\n",
" return {\"score\": int(score), \"key\": \"tool_calls_in_exact_order\"}"
]
},
{
"cell_type": "code",
@@ -355,7 +510,20 @@
]
}
],
"source": ["from langsmith.evaluation import evaluate\n\nexperiment_prefix = f\"custom-agent-{model_tested}\"\nexperiment_results = evaluate(\n predict_custom_agent_local_answer,\n data=dataset_name,\n evaluators=[answer_evaluator, check_trajectory_custom],\n experiment_prefix=experiment_prefix + \"-answer-and-tool-use\",\n num_repetitions=3,\n max_concurrency=1, # Use when running locally\n metadata={\"version\": metadata},\n)"]
"source": [
"from langsmith.evaluation import evaluate\n",
"\n",
"experiment_prefix = f\"custom-agent-{model_tested}\"\n",
"experiment_results = evaluate(\n",
" predict_custom_agent_local_answer,\n",
" data=dataset_name,\n",
" evaluators=[answer_evaluator, check_trajectory_custom],\n",
" experiment_prefix=experiment_prefix + \"-answer-and-tool-use\",\n",
" num_repetitions=3,\n",
" max_concurrency=1, # Use when running locally\n",
" metadata={\"version\": metadata},\n",
")"
]
},
{
"attachments": {
@@ -382,7 +550,9 @@
"id": "79295798-0181-417e-abad-11dddb6ff05e",
"metadata": {},
"outputs": [],
"source": [""]
"source": [
""
]
}
],
"metadata": {
File diff suppressed because one or more lines are too long