File size: 2,996 Bytes
a556e45
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Run GAIA questions concurrently. Each question gets its own agent instance.

Wall-clock is set by the slowest single question rather than the sum of all of
them, so 15 questions take about as long as the worst one β€” a few minutes instead
of an hour.

Usage:
    python run_parallel.py              # all questions without attachments
    python run_parallel.py --workers 8  # tune concurrency
    python run_parallel.py --all        # include file-attachment questions
"""

import json
import sys
import threading
import time
from concurrent.futures import ThreadPoolExecutor, as_completed
from pathlib import Path

from agent import GaiaAgent

HERE = Path(__file__).parent
QUESTIONS = HERE / "gaia_questions.json"
ANSWERS = HERE / "answers.json"

_print_lock = threading.Lock()


def log(msg: str):
    with _print_lock:
        print(msg, flush=True)


def answer_one(index: int, total: int, q: dict) -> dict:
    """Runs one question on its own agent. Never raises β€” a failure returns ''."""
    label = q["question"][:60].replace("\n", " ")
    start = time.time()
    try:
        agent = GaiaAgent()
        answer = agent(
            q["question"], task_id=q["task_id"], file_name=q.get("file_name", "")
        )
    except Exception as e:
        log(f"[{index}/{total}] FAILED {type(e).__name__}: {label}")
        answer = ""

    log(f"[{index}/{total}] {time.time() - start:5.0f}s  {answer!r:<40} | {label}")
    return {"task_id": q["task_id"], "submitted_answer": answer}


def main():
    args = sys.argv[1:]
    include_files = "--all" in args
    workers = 6
    if "--workers" in args:
        workers = int(args[args.index("--workers") + 1])

    questions = json.loads(QUESTIONS.read_text())
    if not include_files:
        questions = [q for q in questions if not q.get("file_name")]

    # --only 3,5,6 keeps just those 1-indexed questions, in the order given.
    # Free-tier token budgets are small, so spending them on the questions most
    # likely to land beats spreading them evenly across ones that cannot.
    if "--only" in args:
        picks = [int(n) for n in args[args.index("--only") + 1].split(",")]
        questions = [questions[i - 1] for i in picks if 1 <= i <= len(questions)]

    total = len(questions)
    log(f"Running {total} questions with {workers} workers\n")
    started = time.time()

    results = []
    with ThreadPoolExecutor(max_workers=workers) as pool:
        futures = {
            pool.submit(answer_one, i, total, q): q
            for i, q in enumerate(questions, 1)
        }
        for fut in as_completed(futures):
            results.append(fut.result())
            ANSWERS.write_text(json.dumps(results, indent=2))

    answered = sum(1 for r in results if r["submitted_answer"])
    log(f"\n{'=' * 70}")
    log(f"Answered {answered}/{total} in {time.time() - started:.0f}s β†’ {ANSWERS.name}")
    log("Verify by hand before submitting β€” GAIA grades on exact match.")


if __name__ == "__main__":
    main()