Spaces:
Sleeping
Sleeping
File size: 7,508 Bytes
00b591a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 | import os
import time
import requests
import re
import json
from datasets import load_dataset
from dotenv import load_dotenv
from agent import agent_executor
from langchain_google_genai.chat_models import ChatGoogleGenerativeAIError
# Load environment variables
load_dotenv()
RESULTS_FILE = "gaia_results.json"
def normalize_answer(s: str) -> str:
"""Normalizes variation in final outputs to guarantee fair string matching."""
if not s:
return ""
s = str(s).strip().lower()
# Strip common conversational wrapper formats and script scene elements
s = re.sub(r'^(final answer:|answer:|the answer is\s*|result:|int\.\s*|ext\.\s*)\s*', '', s)
# Remove trailing period, comma, or punctuation marks
s = s.rstrip('.,!?;:')
# Remove internal hyphen variations and extra spaces
s = s.replace('-', ' ')
s = re.sub(r'\s+', ' ', s)
# Remove common unit suffixes
s = re.sub(r'\s*(hours|meters|m\^3|shares|dollars|\$|usd|thousand hours)\s*$', '', s)
# Standardize whole floating strings (e.g., convert "20.0" -> "20")
if re.match(r'^\d+\.0+$', s):
s = s.split('.')[0]
return s.strip("'\"").strip()
def run_gaia_agent_with_retry(question: str, max_retries: int = 5) -> str:
"""Invokes the agent loop using backoff, waiting 65s on token rate limits to clear per-minute quotas."""
delay = 65 # 65 seconds guarantees the 1-minute token bucket resets completely
for attempt in range(max_retries):
try:
response = agent_executor.invoke({"messages": [("user", question)]})
last_msg = response["messages"][-1]
if isinstance(last_msg.content, list) and len(last_msg.content) > 0:
raw_answer = last_msg.content[0].get("text", "")
else:
raw_answer = last_msg.content
return str(raw_answer).strip()
except Exception as e:
error_msg = str(e).upper()
is_quota_issue = (
"429" in error_msg or
"RESOURCE_EXHAUSTED" in error_msg or
"ChatGoogleGenerativeAIError" in str(type(e))
)
if is_quota_issue:
print(f" [Rate/Token Limit hit] Waiting {delay}s for token bucket reset (attempt {attempt + 1}/{max_retries})...")
time.sleep(delay)
else:
print(f"Agent execution encountered an unhandled error: {e}")
return "ERROR"
print("\n[Quota/Token Limit Reached] Exceeded max retries for task.")
return "QUOTA_EXHAUSTED"
print("\n[Quota/Token Limit Reached] API key has fully exhausted its request volume limits.")
return "QUOTA_EXHAUSTED"
def download_gaia_file(file_name: str, task_id: str):
"""Downloads target document artifacts from the HF Hub repo."""
if not file_name:
return None
url = f"https://huggingface.co/datasets/gaia-benchmark/GAIA/resolve/main/2023/validation/{file_name}"
try:
print(f"-> Downloading evaluation file resource: {file_name}")
response = requests.get(url, timeout=20)
if response.status_code == 200:
with open(file_name, "wb") as f:
f.write(response.content)
return os.path.abspath(file_name)
except Exception as e:
print(f"Failed to fetch task file {file_name}: {e}")
return None
if __name__ == "__main__":
print("=== Starting Auto-Detecting GAIA Evaluation Loop ===")
print("Loading dataset from Hugging Face...")
dataset = load_dataset("gaia-benchmark/GAIA", "2023_all", split="validation")
level_1_tasks = []
for row in dataset:
task_level = row.get("level") if "level" in row else row.get("Level")
if str(task_level) == "1":
level_1_tasks.append(row)
print(f"Found {len(level_1_tasks)} Level 1 tasks for evaluation.")
# Load existing execution progress tracking history from local disk
progress = {}
if os.path.exists(RESULTS_FILE):
with open(RESULTS_FILE, "r") as f:
progress = json.load(f)
print(f"Loaded existing progress. {len(progress)} tasks already processed.\n")
else:
print("Starting a fresh evaluation session.\n")
correct_count = sum(1 for status in progress.values() if status == "CORRECT")
processed_this_session = 0
for task in level_1_tasks:
task_id = task['task_id']
# Skip evaluating this task entry if it matches history checkpoints
if task_id in progress:
continue
processed_this_session += 1
print(f"\n--- Session Task {processed_this_session} (ID: {task_id}) ---")
prompt = task["Question"]
file_name = task.get("file_name", "")
expected_answer = str(task["Final answer"]).strip()
local_file_path = download_gaia_file(file_name, task_id)
if local_file_path:
prompt += (
f"\n\n[System Context Note]: An assignment reference asset data file has been securely downloaded "
f"to your local workspace environment filesystem paths at: '{local_file_path}'. "
f"Use your specialized data tools (`read_local_pdf`, `inspect_excel_sheets`, or `execute_python_code`) "
f"to inspect this file to gather information required to compute the exact answer."
)
print(f"Question: {task['Question']}")
# Run LangGraph Agent Pipeline Loop execution
agent_output = run_gaia_agent_with_retry(prompt)
# If the key is completely dead, break out immediately without counting this task as incorrect
if agent_output == "QUOTA_EXHAUSTED":
print("\nExiting current session cleanly. Please update your GEMINI_API_KEY in the .env file before restarting.")
break
norm_agent = normalize_answer(agent_output)
norm_expected = normalize_answer(expected_answer)
print(f"Raw Agent Output: '{agent_output}' (Normalized: '{norm_agent}')")
print(f"Ground Truth: '{expected_answer}' (Normalized: '{norm_expected}')")
if norm_agent == norm_expected and norm_agent != "error":
print(" Result: MATCH (Correct)")
progress[task_id] = "CORRECT"
correct_count += 1
else:
print(" Result: MISMATCH (Incorrect)")
progress[task_id] = "INCORRECT"
# Instantly dump checkpoint update state array out to disk storage
with open(RESULTS_FILE, "w") as f:
json.dump(progress, f, indent=4)
print("Cooldown window active. Waiting 15 seconds...")
time.sleep(15)
print("-" * 50)
total_processed = len(progress)
print(f"\n=== Cumulative Evaluation Progress Summary ===")
print(f"Total Unique Tasks Processed So Far: {total_processed} / {len(level_1_tasks)}")
if total_processed > 0:
print(f"Current Cumulative Score: {correct_count}/{total_processed} ({(correct_count/total_processed)*100:.2f}%)") |