gouatm's picture
Upload agent.py
c594c6a verified
Raw History Blame Contribute Delete
19.1 kB
"""
GAIA Level-1 agent for the Hugging Face Agents Course final assignment.
Design
------
* smolagents CodeAgent: the LLM writes Python that calls tools, so pure-logic
questions (reversed text, operation tables, botany sorting, running an
attached .py file) are solved directly in code.
* Tools: web search, page fetch (markdown), Wikipedia, YouTube transcript,
audio transcription, Excel/CSV loading, image description (vision model).
* Strict system prompt + post-processing so the returned string is ONLY the
answer (the grader does exact match and forbids the text "FINAL ANSWER").
Model selection (environment variables, checked in this order):
MODEL_ID e.g. "anthropic/claude-sonnet-4-5", "gpt-4o",
"gemini/gemini-2.0-flash" -> used through LiteLLM
(needs ANTHROPIC_API_KEY / OPENAI_API_KEY / GEMINI_API_KEY)
HF_MODEL_ID e.g. "Qwen/Qwen2.5-Coder-32B-Instruct" -> HF Inference API
(needs HF_TOKEN); this is the default if nothing else is set.
"""
from __future__ import annotations
import io
import os
import re
import subprocess
import sys
import tempfile
from pathlib import Path
import requests
from smolagents import (
CodeAgent,
DuckDuckGoSearchTool,
InferenceClientModel,
LiteLLMModel,
Tool,
WikipediaSearchTool,
tool,
)
BROWSER_HEADERS = {
"User-Agent": (
"Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 "
"(KHTML, like Gecko) Chrome/128.0.0.0 Safari/537.36"
),
"Accept": "text/html,application/xhtml+xml,application/xml;q=0.9,*/*;q=0.8",
"Accept-Language": "en-US,en;q=0.9",
}
DEFAULT_API_URL = "https://agents-course-unit4-scoring.hf.space"
FILES_DIR = Path(tempfile.gettempdir()) / "gaia_files"
FILES_DIR.mkdir(parents=True, exist_ok=True)
# --------------------------------------------------------------------------- #
# Tools
# --------------------------------------------------------------------------- #
@tool
def youtube_transcript(url: str) -> str:
"""Return the spoken transcript (with timestamps) of a YouTube video.
Use this for questions about what is said in a video.
Args:
url: Full YouTube URL or the 11-character video id.
"""
from youtube_transcript_api import YouTubeTranscriptApi
m = re.search(r"(?:v=|youtu\.be/|shorts/)([A-Za-z0-9_-]{11})", url)
video_id = m.group(1) if m else url.strip()
try:
api = YouTubeTranscriptApi()
fetched = api.fetch(video_id, languages=["en", "en-US", "en-GB"])
lines = [f"[{int(s.start)//60:02d}:{int(s.start)%60:02d}] {s.text}" for s in fetched]
except Exception as e: # noqa: BLE001
try: # fall back to any available language
api = YouTubeTranscriptApi()
tl = api.list(video_id)
t = next(iter(tl)).fetch()
lines = [f"[{int(s.start)//60:02d}:{int(s.start)%60:02d}] {s.text}" for s in t]
except Exception as e2: # noqa: BLE001
return f"Transcript unavailable ({e}; {e2}). Try web_search for a description of the video."
return "\n".join(lines)[:20000]
@tool
def youtube_video_info(url: str) -> str:
"""Return title, description, uploader and other metadata of a YouTube video
(no download). Useful when the transcript is missing or the question is
about the visual content, which may be described in the title/description.
Args:
url: Full YouTube URL.
"""
try:
import yt_dlp
with yt_dlp.YoutubeDL({"quiet": True, "skip_download": True}) as ydl:
info = ydl.extract_info(url, download=False)
keys = ["title", "uploader", "upload_date", "duration", "description", "tags", "categories"]
return "\n".join(f"{k}: {info.get(k)}" for k in keys)[:6000]
except Exception as e: # noqa: BLE001
return f"Could not fetch video info: {e}"
@tool
def transcribe_audio(file_path: str) -> str:
"""Transcribe speech in an audio file (mp3, wav, m4a ...) to text.
Args:
file_path: Local path to the audio file.
"""
# 1) OpenAI Whisper API if a key is present (fast, accurate)
if os.getenv("OPENAI_API_KEY"):
try:
from openai import OpenAI
client = OpenAI()
with open(file_path, "rb") as f:
r = client.audio.transcriptions.create(model="whisper-1", file=f)
return r.text
except Exception as e: # noqa: BLE001
print(f"OpenAI transcription failed: {e}")
# 2) local faster-whisper (CPU) as a fallback
try:
from faster_whisper import WhisperModel
model = WhisperModel("base.en", device="cpu", compute_type="int8")
segments, _ = model.transcribe(file_path, beam_size=5)
return " ".join(s.text.strip() for s in segments)
except Exception as e: # noqa: BLE001
return f"Transcription failed: {e}"
@tool
def read_spreadsheet(file_path: str) -> str:
"""Load an Excel (.xlsx/.xls) or CSV file and return every sheet as text
(header + all rows) so the data can be reasoned about or re-loaded with pandas.
Args:
file_path: Local path to the spreadsheet.
"""
import pandas as pd
if file_path.lower().endswith(".csv"):
sheets = {"csv": pd.read_csv(file_path)}
else:
sheets = pd.read_excel(file_path, sheet_name=None)
out = []
for name, df in sheets.items():
out.append(f"### sheet: {name} shape={df.shape}\ncolumns={list(df.columns)}\n")
out.append(df.to_string(max_rows=200))
return "\n".join(out)[:30000]
@tool
def run_python_file(file_path: str) -> str:
"""Execute an attached Python script in a subprocess and return its stdout/stderr.
Args:
file_path: Local path to the .py file.
"""
try:
r = subprocess.run(
[sys.executable, file_path], capture_output=True, text=True, timeout=120
)
return f"STDOUT:\n{r.stdout}\nSTDERR:\n{r.stderr}"
except subprocess.TimeoutExpired:
return "Script timed out after 120 s."
@tool
def read_text_file(file_path: str) -> str:
"""Return the raw contents of a text-like file (txt, py, json, md, csv ...).
Args:
file_path: Local path to the file.
"""
return Path(file_path).read_text(errors="replace")[:30000]
@tool
def visit_webpage(url: str) -> str:
"""Visit a webpage and return its content converted to markdown, truncated to a
reasonable length. Use this to read the actual content of a search result,
article, or any other page (not for files like images/audio/xlsx - use the
dedicated tools for those).
Args:
url: The URL to visit.
"""
try:
import markdownify
resp = requests.get(url, headers=BROWSER_HEADERS, timeout=20)
resp.raise_for_status()
md = markdownify.markdownify(resp.text, heading_style="ATX")
md = re.sub(r"\n{3,}", "\n\n", md).strip()
return md[:40000]
except requests.exceptions.HTTPError as e:
return f"Error fetching the webpage: {e}"
except Exception as e: # noqa: BLE001
return f"Error fetching the webpage: {e}"
class DescribeImageTool(Tool):
name = "describe_image"
description = (
"Send an image to a vision language model together with a question and get a "
"detailed textual answer. Use it for any question about an attached image "
"(e.g. a chess board: ask for the full piece placement in FEN and the side to move)."
)
inputs = {
"file_path": {"type": "string", "description": "Local path of the image file."},
"question": {"type": "string", "description": "What to analyse / describe in the image."},
}
output_type = "string"
def __init__(self, model):
super().__init__()
self.model = model
def forward(self, file_path: str, question: str) -> str:
from PIL import Image
from smolagents.models import ChatMessage, MessageRole
img = Image.open(file_path)
msg = ChatMessage(
role=MessageRole.USER,
content=[{"type": "text", "text": question}, {"type": "image", "image": img}],
)
return self.model([msg]).content
@tool
def chess_best_move(fen: str) -> str:
"""Return the best move (SAN) for the side to move in a chess position, using
the Stockfish engine if installed, otherwise a python-chess heuristic search.
Args:
fen: The position in FEN notation, e.g. 'rnbqkbnr/pppppppp/8/8/8/8/PPPPPPPP/RNBQKBNR w KQkq - 0 1'.
"""
import chess
board = chess.Board(fen)
import shutil
engine_path = shutil.which("stockfish")
if engine_path:
import chess.engine
with chess.engine.SimpleEngine.popen_uci(engine_path) as eng:
info = eng.analyse(board, chess.engine.Limit(time=5.0), multipv=3)
lines = []
for i in info:
mv = i["pv"][0]
lines.append(f"{board.san(mv)} score={i['score'].pov(board.turn)}")
return "Candidate moves (best first):\n" + "\n".join(lines)
# Fallback: 3-ply material search, prefers checkmates
values = {chess.PAWN: 1, chess.KNIGHT: 3, chess.BISHOP: 3, chess.ROOK: 5, chess.QUEEN: 9, chess.KING: 0}
def evaluate(b: chess.Board) -> float:
if b.is_checkmate():
return -1000 if b.turn == board.turn else 1000
s = 0
for p, v in values.items():
s += v * (len(b.pieces(p, board.turn)) - len(b.pieces(p, not board.turn)))
return s
def search(b: chess.Board, depth: int, alpha: float, beta: float, maximizing: bool) -> float:
if depth == 0 or b.is_game_over():
return evaluate(b)
best = -1e9 if maximizing else 1e9
for mv in b.legal_moves:
b.push(mv)
v = search(b, depth - 1, alpha, beta, not maximizing)
b.pop()
if maximizing:
best = max(best, v); alpha = max(alpha, v)
else:
best = min(best, v); beta = min(beta, v)
if beta <= alpha:
break
return best
scored = []
for mv in board.legal_moves:
board.push(mv)
scored.append((search(board, 3, -1e9, 1e9, False), mv))
board.pop()
scored.sort(key=lambda t: -t[0])
return "Candidate moves (best first):\n" + "\n".join(f"{board.san(m)} score={s}" for s, m in scored[:5])
# --------------------------------------------------------------------------- #
# Prompt
# --------------------------------------------------------------------------- #
SYSTEM_INSTRUCTIONS = """
You are a general AI assistant. I will ask you a question. Report your thoughts, and
finish your answer with the following template: FINAL ANSWER: [YOUR FINAL ANSWER].
YOUR FINAL ANSWER should be a number OR as few words as possible OR a comma separated
list of numbers and/or strings. If you are asked for a number, don't use comma to
write your number neither use units such as $ or percent sign unless specified
otherwise. If you are asked for a string, don't use articles, neither abbreviations
(e.g. for cities), and write the digits in plain text unless specified otherwise. If
you are asked for a comma separated list, apply the above rules depending on whether
the element to be put in the list is a number or a string.
This is the exact phrasing GAIA's own authors use to prompt models, because the
grader (question_scorer) does quasi-exact-match with normalization tied to the type
of the ground truth:
- If the ground truth is a number: both sides are stripped of "$", "%", "," and
compared as floats. So skip thousands separators and currency symbols unless the
question explicitly asks for them, and give the requested number of decimals.
- If the ground truth contains "," or ";": it is scored as a list. The two lists are
split on [,;] and MUST have the exact same number of elements (a wrong count is an
automatic zero, not partial credit), compared pairwise in the same order as the
ground truth's likely order (respect any ordering instruction, e.g. "alphabetical"
or "clockwise from 12 o'clock").
- Otherwise: whitespace is stripped and case is ignored, but everything else must
match exactly, so do not add trailing periods, quotes, explanations, or the words
"FINAL ANSWER" inside the value itself.
Method:
1. Read the question carefully and note the exact output format requested
(list order, plural/singular, first name only, units, decimals, IOC code, ...).
2. If a file is attached, open it with the appropriate tool first
(read_spreadsheet, run_python_file, transcribe_audio, describe_image, read_text_file).
3. For factual questions, search the web / Wikipedia and READ the actual page with
visit_webpage; do not rely on memory when a source is available. Cross-check
when sources disagree. Respect date qualifiers ("as of 2022", "July 2023").
4. For puzzles / logic / data, compute with Python instead of guessing.
5. Think step by step, then give the answer.
Final answer rules (very important):
- Call final_answer() with ONLY the answer value itself: no explanation, no the
literal text "FINAL ANSWER", no trailing period, no surrounding quotes.
- Numbers: plain digits, no thousands-separator commas, no currency sign or percent
sign unless the question specifically asks for one; match the requested decimals.
- Strings: as few words as possible, no articles ("a", "the"), no abbreviations
(spell out city/country names in full unless the question asks for an abbreviation
or code, e.g. an IOC code), digits written in plain text (e.g. "seven" not "7")
unless the question is itself asking for a numeral.
- Lists: comma-separated, one space after each comma, exactly as many elements as
the question implies, in the exact order the question asks for; apply the number
or string rules above to each element individually.
"""
# --------------------------------------------------------------------------- #
# Model factory
# --------------------------------------------------------------------------- #
def build_model():
model_id = os.getenv("MODEL_ID")
if model_id:
api_base = os.getenv("API_BASE") # e.g. a public tunnel URL for a local Ollama
print(f"Using LiteLLM model: {model_id}" + (f" via {api_base}" if api_base else ""))
kwargs = {"model_id": model_id, "temperature": 0.0, "max_tokens": 4096}
if api_base:
kwargs["api_base"] = api_base
if model_id.startswith("ollama") and not os.getenv("OLLAMA_API_KEY"):
# LiteLLM's ollama routes don't need a real key, but some paths still
# check for one being set; a harmless placeholder avoids that.
kwargs["api_key"] = "ollama"
if api_base and "ngrok" in api_base:
# ngrok's free tier serves an HTML interstitial warning page to any
# request without this header, which breaks the Ollama API calls with
# a cryptic empty PermissionDeniedError.
kwargs["extra_headers"] = {"ngrok-skip-browser-warning": "true"}
return LiteLLMModel(**kwargs)
hf_model = os.getenv("HF_MODEL_ID", "Qwen/Qwen2.5-Coder-32B-Instruct")
print(f"Using HF Inference model: {hf_model}")
return InferenceClientModel(model_id=hf_model, token=os.getenv("HF_TOKEN"), temperature=0.0)
def build_vision_model():
"""A model that can look at images (separate env var so a cheap text model can be
used for the main loop)."""
vid = os.getenv("VISION_MODEL_ID") or os.getenv("MODEL_ID")
if vid:
return LiteLLMModel(model_id=vid, temperature=0.0, max_tokens=2048)
return InferenceClientModel(
model_id=os.getenv("HF_VISION_MODEL_ID", "Qwen/Qwen2.5-VL-72B-Instruct"),
token=os.getenv("HF_TOKEN"),
)
# --------------------------------------------------------------------------- #
# Answer cleanup
# --------------------------------------------------------------------------- #
def clean_answer(ans) -> str:
s = str(ans).strip()
s = re.sub(r"(?i)^\s*final answer\s*:?\s*", "", s)
for _ in range(3): # peel quotes / trailing periods in any order
s = s.strip().strip('"').strip("'").strip("`").strip()
if s.endswith("."):
s = s[:-1]
s = s.strip()
# normalise "a,b ,c" -> "a, b, c"
if "," in s:
s = ", ".join(p.strip() for p in s.split(","))
return s
# --------------------------------------------------------------------------- #
# Agent wrapper
# --------------------------------------------------------------------------- #
class GaiaAgent:
def __init__(self, api_url: str = DEFAULT_API_URL):
self.api_url = api_url
self.model = build_model()
vision = build_vision_model()
tools = [
DuckDuckGoSearchTool(max_results=8),
visit_webpage,
WikipediaSearchTool(user_agent="GaiaCourseAgent/1.0 (https://huggingface.co/spaces; student project)", content_type="text", extract_format="WIKI"),
youtube_transcript,
youtube_video_info,
transcribe_audio,
read_spreadsheet,
run_python_file,
read_text_file,
DescribeImageTool(vision),
chess_best_move,
]
self.agent = CodeAgent(
tools=tools,
model=self.model,
additional_authorized_imports=[
"pandas", "numpy", "json", "re", "itertools", "collections", "math",
"datetime", "statistics", "csv", "bs4", "requests", "chess",
],
max_steps=20,
planning_interval=4,
verbosity_level=1,
instructions=SYSTEM_INSTRUCTIONS,
)
print("GaiaAgent initialized.")
# ---- file handling -------------------------------------------------- #
def fetch_file(self, task_id: str, file_name: str) -> str | None:
if not file_name:
return None
dest = FILES_DIR / file_name
if dest.exists():
return str(dest)
try:
r = requests.get(f"{self.api_url}/files/{task_id}", timeout=60)
r.raise_for_status()
dest.write_bytes(r.content)
return str(dest)
except Exception as e: # noqa: BLE001
print(f"Could not download file for {task_id}: {e}")
return None
# ---- main entry ----------------------------------------------------- #
def __call__(self, question: str, task_id: str | None = None, file_name: str = "") -> str:
task = question
path = self.fetch_file(task_id, file_name) if task_id else None
if path:
task += f"\n\n[The file mentioned in the question is available locally at: {path}]"
try:
result = self.agent.run(task)
except Exception as e: # noqa: BLE001
print(f"Agent error: {e}")
return "unknown"
return clean_answer(result)