| from smolagents import Tool |
| from langchain_community.retrievers import BM25Retriever |
| from langchain.docstore.document import Document |
| from tools import CrashInfoRetrieverTool |
| import datasets |
| import pandas as pd |
| import requests |
|
|
| class GuestInfoRetrieverTool(Tool): |
| name = "guest_info_retriever" |
| description = "Retrieves detailed information about gala guests based on their name or relation." |
| inputs = { |
| "query": { |
| "type": "string", |
| "description": "The name or relation of the guest you want information about." |
| } |
| } |
| output_type = "string" |
|
|
| def __init__(self, docs): |
| self.is_initialized = False |
| self.retriever = BM25Retriever.from_documents(docs) |
| |
|
|
| def forward(self, query: str): |
| results = self.retriever.get_relevant_documents(query) |
| if results: |
| return "\n\n".join([doc.page_content for doc in results[:3]]) |
| else: |
| return "No matching guest information found." |
|
|
|
|
| def load_guest_dataset(): |
| |
| guest_dataset = datasets.load_dataset("agents-course/unit3-invitees", split="train") |
|
|
| |
| docs = [ |
| Document( |
| page_content="\n".join([ |
| f"Name: {guest['name']}", |
| f"Relation: {guest['relation']}", |
| f"Description: {guest['description']}", |
| f"Email: {guest['email']}" |
| ]), |
| metadata={"name": guest["name"]} |
| ) |
| for guest in guest_dataset |
| ] |
|
|
| |
| return GuestInfoRetrieverTool(docs) |
|
|
| def load_crash_data(): |
| |
| url = "https://phl.carto.com/api/v2/sql?filename=fatal_crashes&format=csv&skipfields=cartodb_id,the_geom,the_geom_webmercator&q=SELECT%20*,%20ST_Y(the_geom)%20AS%20lat,%20ST_X(the_geom)%20AS%20lng%20FROM%20fatal_crashes" |
| response = requests.get(url) |
| with open("fatal_crashes.csv", "wb") as f: |
| f.write(response.content) |
|
|
| df = pd.read_csv("fatal_crashes.csv") |
|
|
| |
| docs = [] |
| for _, row in df.iterrows(): |
| content = "\n".join([ |
| f"Year: {row.get('year', 'N/A')}", |
| f"Street: {row.get('primary_st', 'N/A')}", |
| f"Age: {row.get('age', 'N/A')}", |
| f"Arrest: {row.get('arrest_yes', 'N/A')}", |
| f"Primary Vehicle: {row.get('veh1', 'N/A')}", |
| f"Crash Type: {row.get('veh2', 'N/A')}", |
| f"Outcome: {row.get('investigat', 'N/A')}", |
| f"RecordID: {row.get('objectid', 'N/A')}" |
| ]) |
| docs.append(Document(page_content=content, metadata={"id": row.get("id", "unknown")})) |
|
|
| |
| return CrashInfoRetrieverTool(docs) |
|
|
|
|