Spaces:
Running on Zero
Running on Zero
Download agent.py from gouatm/Final_Assignment_Template: direct link, hf CLI and curl.
- Browser
- Download file 19.1 kB
-
https://huggingface.co/spaces/gouatm/Final_Assignment_Template/resolve/main/agent.py
- Command line
-
hf download hf://spaces/gouatm/Final_Assignment_Template/agent.py
-
curl -L -o agent.py https://huggingface.co/spaces/gouatm/Final_Assignment_Template/resolve/main/agent.py
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 | |
| # --------------------------------------------------------------------------- # | |
| 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] | |
| 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}" | |
| 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}" | |
| 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] | |
| 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." | |
| 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] | |
| 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 | |
| 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) | |