{
"cells": [
{
"attachments": {},
"cell_type": "markdown",
"id": "1c0beb97",
"metadata": {},
"source": [
""
]
},
{
"attachments": {},
"cell_type": "markdown",
"id": "8cd3f128-866a-4857-a00a-df19f926c952",
"metadata": {},
"source": [
"# Evaporate Demo\n",
"\n",
"This demo shows how you can extract DataFrame from raw text using the Evaporate paper (Arora et al.): https://arxiv.org/abs/2304.09433.\n",
"\n",
"The inspiration is to first \"fit\" on a set of training text. The fitting process uses the LLM to generate a set of parsing functions from the text.\n",
"These fitted functions are then applied to text during inference time."
]
},
{
"attachments": {},
"cell_type": "markdown",
"id": "7c152e41",
"metadata": {},
"source": [
"If you're opening this Notebook on colab, you will probably need to install LlamaIndex 🦙."
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "b6cc721a",
"metadata": {},
"outputs": [],
"source": [
"%pip install llama-index-llms-openai\n",
"%pip install llama-index-program-evaporate"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "3096cbad",
"metadata": {},
"outputs": [],
"source": [
"!pip install llama-index"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "db7210f2-f19d-4112-ab72-ddb3afe282f7",
"metadata": {},
"outputs": [],
"source": [
"%load_ext autoreload\n",
"%autoreload 2"
]
},
{
"attachments": {},
"cell_type": "markdown",
"id": "da19d340-57b5-439f-9cb1-5ba9576ec304",
"metadata": {},
"source": [
"## Use `DFEvaporateProgram` \n",
"\n",
"The `DFEvaporateProgram` will extract a 2D dataframe from a set of datapoints given a set of fields, and some training data to \"fit\" some functions on."
]
},
{
"attachments": {},
"cell_type": "markdown",
"id": "a299cad8-af81-4974-a3de-ed43877d3490",
"metadata": {},
"source": [
"### Load data\n",
"\n",
"Here we load a set of cities from Wikipedia."
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "daf434f6-3b27-4805-9de8-8fc92d7d776b",
"metadata": {},
"outputs": [],
"source": [
"wiki_titles = [\"Toronto\", \"Seattle\", \"Chicago\", \"Boston\", \"Houston\"]"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "8438168c-3b1b-425e-98b0-2c67a8a58a5f",
"metadata": {},
"outputs": [],
"source": [
"from pathlib import Path\n",
"\n",
"import requests\n",
"\n",
"for title in wiki_titles:\n",
" response = requests.get(\n",
" \"https://en.wikipedia.org/w/api.php\",\n",
" params={\n",
" \"action\": \"query\",\n",
" \"format\": \"json\",\n",
" \"titles\": title,\n",
" \"prop\": \"extracts\",\n",
" # 'exintro': True,\n",
" \"explaintext\": True,\n",
" },\n",
" ).json()\n",
" page = next(iter(response[\"query\"][\"pages\"].values()))\n",
" wiki_text = page[\"extract\"]\n",
"\n",
" data_path = Path(\"data\")\n",
" if not data_path.exists():\n",
" Path.mkdir(data_path)\n",
"\n",
" with open(data_path / f\"{title}.txt\", \"w\") as fp:\n",
" fp.write(wiki_text)"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "c01dbcb8-5ea1-4e76-b5de-ea5ebe4f0392",
"metadata": {},
"outputs": [],
"source": [
"from llama_index.core import SimpleDirectoryReader\n",
"\n",
"# Load all wiki documents\n",
"city_docs = {}\n",
"for wiki_title in wiki_titles:\n",
" city_docs[wiki_title] = SimpleDirectoryReader(\n",
" input_files=[f\"data/{wiki_title}.txt\"]\n",
" ).load_data()"
]
},
{
"attachments": {},
"cell_type": "markdown",
"id": "e7310883-2aeb-4a4d-b101-b3279e670ea8",
"metadata": {},
"source": [
"### Parse Data"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "b8e98279-b4c4-41ec-b696-13e6a6f841a4",
"metadata": {},
"outputs": [],
"source": [
"from llama_index.llms.openai import OpenAI\n",
"from llama_index.core import Settings\n",
"\n",
"# setup settings\n",
"Settings.llm = OpenAI(temperature=0, model=\"gpt-3.5-turbo\")\n",
"Settings.chunk_size = 512"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "74c6c1c3-b797-45c8-b692-7a6e4bd1898d",
"metadata": {},
"outputs": [],
"source": [
"# get nodes for each document\n",
"city_nodes = {}\n",
"for wiki_title in wiki_titles:\n",
" docs = city_docs[wiki_title]\n",
" nodes = Settings.node_parser.get_nodes_from_documents(docs)\n",
" city_nodes[wiki_title] = nodes"
]
},
{
"attachments": {},
"cell_type": "markdown",
"id": "bb369a78-e634-43f4-805e-52f6ea0f3588",
"metadata": {},
"source": [
"### Running the DFEvaporateProgram\n",
"\n",
"Here we demonstrate how to extract datapoints with our `DFEvaporateProgram`. Given a set of fields, the `DFEvaporateProgram` can first fit functions on a set of training data, and then run extraction over inference data."
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "6c260836",
"metadata": {},
"outputs": [],
"source": [
"from llama_index.program.evaporate import DFEvaporateProgram\n",
"\n",
"# define program\n",
"program = DFEvaporateProgram.from_defaults(\n",
" fields_to_extract=[\"population\"],\n",
")"
]
},
{
"attachments": {},
"cell_type": "markdown",
"id": "c548768e-9d4a-4708-9c84-9266503edf01",
"metadata": {},
"source": [
"### Fitting Functions"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "6c186eb7-116f-4b28-a508-8639cbc86633",
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"{'population': 'def get_population_field(text: str):\\n \"\"\"\\n Function to extract population. \\n \"\"\"\\n \\n # Use regex to find the population field\\n pattern = r\\'(?<=population of )(\\\\d+,?\\\\d*)\\'\\n population_field = re.search(pattern, text).group(1)\\n \\n # Return the population field as a single value\\n return int(population_field.replace(\\',\\', \\'\\'))'}"
]
},
"execution_count": null,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"program.fit_fields(city_nodes[\"Toronto\"][:1])"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "483676a4-4937-40a8-acd9-8fec4a991270",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"def get_population_field(text: str):\n",
" \"\"\"\n",
" Function to extract population. \n",
" \"\"\"\n",
" \n",
" # Use regex to find the population field\n",
" pattern = r'(?<=population of )(\\d+,?\\d*)'\n",
" population_field = re.search(pattern, text).group(1)\n",
" \n",
" # Return the population field as a single value\n",
" return int(population_field.replace(',', ''))\n"
]
}
],
"source": [
"# view extracted function\n",
"print(program.get_function_str(\"population\"))"
]
},
{
"attachments": {},
"cell_type": "markdown",
"id": "508a442c-d7d8-4a27-8add-1d58f1ecc66b",
"metadata": {},
"source": [
"### Run Inference"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "83e38b62-bad0-4154-9597-555a27e976d9",
"metadata": {},
"outputs": [],
"source": [
"seattle_df = program(nodes=city_nodes[\"Seattle\"][:1])"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "dc72f611-da8b-4882-b532-69e46b9589bb",
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"DataFrameRowsOnly(rows=[DataFrameRow(row_values=[749256])])"
]
},
"execution_count": null,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"seattle_df"
]
},
{
"attachments": {},
"cell_type": "markdown",
"id": "9465ba41-8318-40bb-a202-49df6e3c16e3",
"metadata": {},
"source": [
"## Use `MultiValueEvaporateProgram` \n",
"\n",
"In contrast to the `DFEvaporateProgram`, which assumes the output obeys a 2D tabular format (one row per node), the `MultiValueEvaporateProgram` returns a list of `DataFrameRow` objects - each object corresponds to a column, and can contain a variable length of values. This can help if we want to extract multiple values for one field from a given piece of text.\n",
"\n",
"In this example, we use this program to parse gold medal counts."
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "c3d5e9dd-0d20-447b-96b2-a82f8350e430",
"metadata": {},
"outputs": [],
"source": [
"Settings.llm = OpenAI(temperature=0, model=\"gpt-4\")\n",
"Settings.chunk_size = 1024\n",
"Settings.chunk_overlap = 0"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "08b44698-4f7e-4686-9b6e-1b77c341a778",
"metadata": {},
"outputs": [],
"source": [
"from llama_index.core.data_structs import Node\n",
"\n",
"# Olympic total medal counts: https://en.wikipedia.org/wiki/All-time_Olympic_Games_medal_table\n",
"\n",
"train_text = \"\"\"\n",
"
| Team (IOC code)\n", " | \n", "No. Summer\n", " | \n", "No. Winter\n", " | \n", "No. Games\n", " |
|---|---|---|---|
| 9 | \n", "5 | \n", "14\n", " | |
| 9 | \n", "2 | \n", "11\n", " | |
| 12 | \n", "13 | \n", "25\n", " | |
| 10 | \n", "0 | \n", "10\n", " | |
| 11 | \n", "0 | \n", "11\n", " | |
| 9 | \n", "0 | \n", "9\n", " | 10 | \n", "0 | \n", "10\n", " | \n", "
| 13 | \n", "0 | \n", "13\n", " | |
| 12 | \n", "0 | \n", "12\n", " | |
| 10 | \n", "0 | \n", "10\n", " | |
| 15 | \n", "7 | \n", "22\n", " | |
| 8 | \n", "8 | \n", "16\n", " | |
| 10 | \n", "2 | \n", "12\n", " | |
| 6 | \n", "0 | \n", "6\n", " | |
| 10 | \n", "0 | \n", "10\n", " | |
| 7 | \n", "0 | \n", "7\n", " | |
| 11 | \n", "2 | \n", "13\n", " | |
| 11 | \n", "0 | \n", "11\n", " | |
| 13 | \n", "0 | \n", "13\n", " | |
| 7 | \n", "0 | \n", "7\n", " | |
| 13 | \n", "0 | \n", "13\n", " | |
| 11 | \n", "0 | \n", "11\n", " | (.*?) | \\', text)\\n \\n # Return the result as a list\\n return medal_count_field'}" ] }, "execution_count": null, "metadata": {}, "output_type": "execute_result" } ], "source": [ "program.fit_fields(train_nodes[:1])" ] }, { "cell_type": "code", "execution_count": null, "id": "cc32440c-910a-483c-81df-80ae81fedb2d", "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "def get_countries_field(text: str):\n", " \"\"\"\n", " Function to extract countries. \n", " \"\"\"\n", " \n", " # Use regex to extract the countries field\n", " countries_field = re.findall(r'(.*)', text)\n", " \n", " # Return the result as a list\n", " return countries_field\n" ] } ], "source": [ "print(program.get_function_str(\"countries\"))" ] }, { "cell_type": "code", "execution_count": null, "id": "8ed16aa9-8b36-439a-a596-1b90d6775a30", "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "def get_medal_count_field(text: str):\n", " \"\"\"\n", " Function to extract medal_count. \n", " \"\"\"\n", " \n", " # Use regex to extract the medal count field\n", " medal_count_field = re.findall(r'(.*?) | ', text)\n", " \n", " # Return the result as a list\n", " return medal_count_field\n" ] } ], "source": [ "print(program.get_function_str(\"medal_count\"))" ] }, { "cell_type": "code", "execution_count": null, "id": "8f7bae5f-ee4e-4d9f-b551-1986efd317b3", "metadata": {}, "outputs": [], "source": [ "result = program(nodes=infer_nodes[:1])" ] }, { "cell_type": "code", "execution_count": null, "id": "85bc4f9c-9e6c-41da-b6fb-a8b227b3ce67", "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "Countries: ['Bangladesh', '[BIZ]', '[BEN]', 'Bhutan', 'Bolivia', 'Bosnia and Herzegovina', 'British Virgin Islands', '[A]', 'Cambodia', 'Cape Verde', 'Cayman Islands', 'Central African Republic', 'Chad', 'Comoros', 'Republic of the Congo', '[COD]']\n", "\n", "Medal Counts: ['Bangladesh', '[BIZ]', '[BEN]', 'Bhutan', 'Bolivia', 'Bosnia and Herzegovina', 'British Virgin Islands', '[A]', 'Cambodia', 'Cape Verde', 'Cayman Islands', 'Central African Republic', 'Chad', 'Comoros', 'Republic of the Congo', '[COD]']\n", "\n" ] } ], "source": [ "# output countries\n", "print(f\"Countries: {result.columns[0].row_values}\\n\")\n", "# output medal counts\n", "print(f\"Medal Counts: {result.columns[0].row_values}\\n\")" ] }, { "attachments": {}, "cell_type": "markdown", "id": "820768fe-aa23-4999-bcc1-102e6fc817e5", "metadata": {}, "source": [ "## Bonus: Use the underlying `EvaporateExtractor`\n", "\n", "The underlying `EvaporateExtractor` offers some additional functionality, e.g. actually helping to identify fields over a set of text.\n", "\n", "Here we show how you can use `identify_fields` to determine relevant fields around a general `topic` field." ] }, { "cell_type": "code", "execution_count": null, "id": "7ff32b4b-a85b-4266-bdf1-7fa492925034", "metadata": {}, "outputs": [], "source": [ "# a list of nodes, one node per city, corresponding to intro paragraph\n", "# city_pop_nodes = []\n", "city_pop_nodes = [city_nodes[\"Toronto\"][0], city_nodes[\"Seattle\"][0]]" ] }, { "cell_type": "code", "execution_count": null, "id": "dc96646f-ac7e-407f-87dd-c14c8d83aa84", "metadata": {}, "outputs": [], "source": [ "extractor = program.extractor" ] }, { "cell_type": "code", "execution_count": null, "id": "1df3a7df-6d00-4487-b114-f45a6dba4764", "metadata": {}, "outputs": [], "source": [ "# Try with Toronto and Seattle (should extract \"population\")\n", "existing_fields = extractor.identify_fields(\n", " city_pop_nodes, topic=\"population\", fields_top_k=4\n", ")" ] }, { "cell_type": "code", "execution_count": null, "id": "d8a56bb6-3a26-40db-9ca3-8aa9ed4f2c52", "metadata": {}, "outputs": [ { "data": { "text/plain": [ "[\"seattle metropolitan area's population\"]" ] }, "execution_count": null, "metadata": {}, "output_type": "execute_result" } ], "source": [ "existing_fields" ] } ], "metadata": { "kernelspec": { "display_name": "llama_index_v2", "language": "python", "name": "llama_index_v2" }, "language_info": { "codemirror_mode": { "name": "ipython", "version": 3 }, "file_extension": ".py", "mimetype": "text/x-python", "name": "python", "nbconvert_exporter": "python", "pygments_lexer": "ipython3" } }, "nbformat": 4, "nbformat_minor": 5 }