Spaces:
Sleeping
Sleeping
Update app.py
Browse files
app.py
CHANGED
|
@@ -3,107 +3,141 @@ import gradio as gr
|
|
| 3 |
import requests
|
| 4 |
import pandas as pd
|
| 5 |
import time
|
| 6 |
-
import random
|
| 7 |
import re
|
| 8 |
import math
|
| 9 |
-
from typing import Dict, Any, List, Optional
|
| 10 |
|
| 11 |
# (Keep Constants as is)
|
| 12 |
# --- Constants ---
|
| 13 |
DEFAULT_API_URL = "https://agents-course-unit4-scoring.hf.space"
|
| 14 |
|
| 15 |
-
# ---
|
| 16 |
|
| 17 |
-
|
| 18 |
"""
|
| 19 |
-
|
| 20 |
-
Supported providers via env LLM_PROVIDER:
|
| 21 |
-
- openai (OPENAI_API_KEY, model via OPENAI_MODEL, base via OPENAI_BASE_URL)
|
| 22 |
-
- together (TOGETHER_API_KEY, model via TOGETHER_MODEL, base via TOGETHER_BASE_URL)
|
| 23 |
-
- openrouter (OPENROUTER_API_KEY, model via OPENROUTER_MODEL, base fixed)
|
| 24 |
-
Defaults pick small reasoning-capable models.
|
| 25 |
"""
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 26 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 27 |
def __init__(self):
|
| 28 |
-
self.provider =
|
|
|
|
|
|
|
|
|
|
| 29 |
|
| 30 |
-
|
| 31 |
-
|
| 32 |
-
|
| 33 |
-
|
| 34 |
-
|
| 35 |
-
self.model = os.getenv("OPENAI_MODEL", "gpt-4o-mini")
|
| 36 |
self.base_url = os.getenv("OPENAI_BASE_URL", "https://api.openai.com/v1")
|
| 37 |
-
self.
|
| 38 |
-
|
| 39 |
-
|
| 40 |
-
|
| 41 |
-
|
| 42 |
-
elif self.provider == "together":
|
| 43 |
-
self.api_key = os.getenv("TOGETHER_API_KEY")
|
| 44 |
-
if not self.api_key:
|
| 45 |
-
raise ValueError("TOGETHER_API_KEY is required when LLM_PROVIDER=together.")
|
| 46 |
-
self.model = os.getenv("TOGETHER_MODEL", "meta-llama/Meta-Llama-3.1-8B-Instruct-Turbo")
|
| 47 |
-
self.base_url = os.getenv("TOGETHER_BASE_URL", "https://api.together.xyz/v1")
|
| 48 |
-
self.headers = {
|
| 49 |
-
"Authorization": f"Bearer {self.api_key}",
|
| 50 |
-
"Content-Type": "application/json",
|
| 51 |
-
}
|
| 52 |
-
|
| 53 |
-
elif self.provider == "openrouter":
|
| 54 |
-
self.api_key = os.getenv("OPENROUTER_API_KEY")
|
| 55 |
-
if not self.api_key:
|
| 56 |
-
raise ValueError("OPENROUTER_API_KEY is required when LLM_PROVIDER=openrouter.")
|
| 57 |
-
# Strong, cost-effective default; change as desired
|
| 58 |
-
self.model = os.getenv("OPENROUTER_MODEL", "meta-llama/llama-3.1-8b-instruct:free")
|
| 59 |
self.base_url = "https://openrouter.ai/api/v1"
|
| 60 |
-
self.
|
| 61 |
-
|
| 62 |
-
"Content-Type": "application/json",
|
| 63 |
-
}
|
| 64 |
-
|
| 65 |
else:
|
| 66 |
-
|
| 67 |
|
| 68 |
-
|
| 69 |
-
self.
|
| 70 |
-
self.base_backoff = float(os.getenv("BASE_BACKOFF", "1.0")) # seconds
|
| 71 |
-
self.max_backoff = float(os.getenv("MAX_BACKOFF", "30.0")) # cap
|
| 72 |
|
| 73 |
-
def
|
| 74 |
-
|
| 75 |
-
|
| 76 |
-
for attempt in range(1, self.max_retries + 1):
|
| 77 |
-
try:
|
| 78 |
-
resp = requests.post(url, headers=headers, json=payload, timeout=timeout)
|
| 79 |
-
if resp.status_code in (429, 500, 502, 503, 504):
|
| 80 |
-
retry_after = resp.headers.get("Retry-After")
|
| 81 |
-
if retry_after:
|
| 82 |
-
try:
|
| 83 |
-
delay = max(delay, float(retry_after))
|
| 84 |
-
except Exception:
|
| 85 |
-
pass
|
| 86 |
-
jitter = random.uniform(0, delay * 0.25)
|
| 87 |
-
wait_s = min(delay + jitter, self.max_backoff)
|
| 88 |
-
print(f"⏳ {self.provider} {resp.status_code}. Attempt {attempt}/{self.max_retries}. Waiting {wait_s:.2f}s...")
|
| 89 |
-
time.sleep(wait_s)
|
| 90 |
-
delay = min(delay * 1.8, self.max_backoff)
|
| 91 |
-
continue
|
| 92 |
-
return resp
|
| 93 |
-
except requests.exceptions.RequestException as e:
|
| 94 |
-
last_exc = e
|
| 95 |
-
jitter = random.uniform(0, delay * 0.25)
|
| 96 |
-
wait_s = min(delay + jitter, self.max_backoff)
|
| 97 |
-
print(f"⚠️ Network error on attempt {attempt}/{self.max_retries}: {e}. Waiting {wait_s:.2f}s...")
|
| 98 |
-
time.sleep(wait_s)
|
| 99 |
-
delay = min(delay * 1.8, self.max_backoff)
|
| 100 |
-
if last_exc:
|
| 101 |
-
raise last_exc
|
| 102 |
-
return requests.post(url, headers=headers, json=payload, timeout=timeout)
|
| 103 |
-
|
| 104 |
-
def chat(self, system_prompt: str, user_prompt: str, max_tokens: int = 512, temperature: float = 0.6, top_p: float = 0.9) -> str:
|
| 105 |
url = f"{self.base_url}/chat/completions"
|
| 106 |
-
payload
|
| 107 |
"model": self.model,
|
| 108 |
"messages": [
|
| 109 |
{"role": "system", "content": system_prompt},
|
|
@@ -111,153 +145,82 @@ class LLMClient:
|
|
| 111 |
],
|
| 112 |
"max_tokens": max_tokens,
|
| 113 |
"temperature": temperature,
|
| 114 |
-
"top_p":
|
| 115 |
}
|
| 116 |
-
|
| 117 |
-
|
| 118 |
-
|
| 119 |
-
|
| 120 |
-
|
| 121 |
-
|
| 122 |
-
|
| 123 |
-
)
|
| 124 |
-
|
| 125 |
-
|
| 126 |
-
|
| 127 |
-
|
| 128 |
-
|
| 129 |
-
|
| 130 |
-
|
| 131 |
-
|
| 132 |
-
|
| 133 |
-
# --- Simple tool-use: safe-ish arithmetic and unit conversions ---
|
| 134 |
-
|
| 135 |
-
def maybe_compute_locally(question: str) -> Optional[str]:
|
| 136 |
-
"""
|
| 137 |
-
Lightweight math helper for common GAIA-style arithmetic/unit questions.
|
| 138 |
-
Avoids eval of arbitrary text; only allows digits, ops, dots, spaces, and parentheses.
|
| 139 |
-
Returns a computed answer string or None to defer to LLM.
|
| 140 |
-
"""
|
| 141 |
-
# Quick patterns: pure arithmetic, percents, simple ratios, sqrt, powers
|
| 142 |
-
text = question.strip().lower()
|
| 143 |
-
|
| 144 |
-
# Extract expression inside "calculate"/"compute"/"what is"
|
| 145 |
-
match = re.search(r"(?:calculate|compute|what is|evaluate)\s*[:\-]?\s*(.+)", text)
|
| 146 |
-
expr = match.group(1).strip() if match else text
|
| 147 |
-
|
| 148 |
-
# Replace common words to operators
|
| 149 |
-
expr = expr.replace("×", "*").replace("÷", "/").replace("^", "**")
|
| 150 |
-
expr = re.sub(r"(\d+)\s+percent", r"(\1/100)", expr)
|
| 151 |
-
|
| 152 |
-
# Allow only safe characters
|
| 153 |
-
if not re.fullmatch(r"[0-9\.\s\+\-\*\/\(\)\%]+", re.sub(r"\*\*", "**", expr)):
|
| 154 |
-
return None
|
| 155 |
-
|
| 156 |
-
# Disallow modulo for now (often not needed)
|
| 157 |
-
if "%" in expr:
|
| 158 |
-
return None
|
| 159 |
-
|
| 160 |
-
try:
|
| 161 |
-
# Use Python arithmetic safely
|
| 162 |
-
result = eval(expr, {"__builtins__": {}}, {"sqrt": math.sqrt, "pow": pow})
|
| 163 |
-
except Exception:
|
| 164 |
-
return None
|
| 165 |
-
|
| 166 |
-
if isinstance(result, float):
|
| 167 |
-
# Round to a reasonable precision
|
| 168 |
-
result = round(result, 6)
|
| 169 |
-
# normalize -0.0
|
| 170 |
-
if result == 0:
|
| 171 |
-
result = 0.0
|
| 172 |
-
return str(result)
|
| 173 |
-
|
| 174 |
-
# --- Basic Agent Definition ---
|
| 175 |
-
# ----- THIS IS WERE YOU CAN BUILD WHAT YOU WANT ------
|
| 176 |
|
| 177 |
class BasicAgent:
|
| 178 |
def __init__(self):
|
| 179 |
-
# Config
|
| 180 |
-
self.rate_limit_s = float(os.getenv("RATE_LIMIT_SECONDS", "1.0"))
|
| 181 |
-
self.max_new_tokens = int(os.getenv("MAX_NEW_TOKENS", "256"))
|
| 182 |
-
self.samples = int(os.getenv("NUM_SAMPLES", "3")) # self-consistency votes
|
| 183 |
-
self.temperature = float(os.getenv("TEMPERATURE", "0.6"))
|
| 184 |
-
self.top_p = float(os.getenv("TOP_P", "0.9"))
|
| 185 |
-
|
| 186 |
-
# System prompt optimized for concise, correct answers
|
| 187 |
-
self.system_prompt = (
|
| 188 |
-
"You are a precise reasoning assistant. Answer with the final result only, "
|
| 189 |
-
"without extra explanations, unless the question explicitly asks for steps or justification. "
|
| 190 |
-
"Prefer exact values, otherwise numeric to a sensible precision. Be factual and avoid speculation."
|
| 191 |
-
)
|
| 192 |
-
|
| 193 |
-
# LLM client (non-HF endpoints)
|
| 194 |
self.client = LLMClient()
|
| 195 |
-
|
|
|
|
| 196 |
|
| 197 |
def _finalize(self, text: str) -> str:
|
| 198 |
-
|
| 199 |
-
t = text.strip()
|
| 200 |
-
t = re.sub(r"^```.*?\n", "", t, flags=re.DOTALL) # remove starting fence if any
|
| 201 |
-
t = t.strip("` \n")
|
| 202 |
-
# Collapse multiple spaces
|
| 203 |
t = re.sub(r"\s+", " ", t)
|
| 204 |
-
return t
|
| 205 |
|
| 206 |
-
def
|
| 207 |
-
|
|
|
|
|
|
|
| 208 |
local = maybe_compute_locally(question)
|
| 209 |
if local is not None:
|
|
|
|
| 210 |
return local
|
| 211 |
|
| 212 |
-
#
|
| 213 |
-
|
| 214 |
-
|
| 215 |
-
|
| 216 |
-
|
| 217 |
-
|
| 218 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 219 |
)
|
| 220 |
-
|
| 221 |
|
| 222 |
-
|
| 223 |
-
|
| 224 |
|
| 225 |
-
#
|
| 226 |
-
|
| 227 |
-
|
| 228 |
-
|
| 229 |
-
|
| 230 |
-
|
| 231 |
-
|
| 232 |
-
|
| 233 |
-
|
| 234 |
-
# Majority vote by normalized string; numeric-aware normalization
|
| 235 |
-
normalized_counts: Dict[str, int] = {}
|
| 236 |
-
mapping: Dict[str, str] = {}
|
| 237 |
-
|
| 238 |
-
def normalize(s: str) -> str:
|
| 239 |
-
s_clean = s.strip().lower()
|
| 240 |
-
# Try extract number
|
| 241 |
-
num = re.findall(r"-?\d+(?:\.\d+)?", s_clean)
|
| 242 |
-
if len(num) == 1 and s_clean.replace(num[0], "").strip() in ["", "units", "unit"]:
|
| 243 |
-
try:
|
| 244 |
-
val = float(num[0])
|
| 245 |
-
return f"{round(val, 6)}"
|
| 246 |
-
except Exception:
|
| 247 |
-
pass
|
| 248 |
-
return s_clean
|
| 249 |
-
|
| 250 |
-
for c in candidates:
|
| 251 |
-
key = normalize(c)
|
| 252 |
-
normalized_counts[key] = normalized_counts.get(key, 0) + 1
|
| 253 |
-
# Keep first original for formatting
|
| 254 |
-
if key not in mapping:
|
| 255 |
-
mapping[key] = c
|
| 256 |
-
|
| 257 |
-
best_key = max(normalized_counts, key=normalized_counts.get)
|
| 258 |
-
final = mapping[best_key]
|
| 259 |
-
print(f"Agent returning answer: {final}")
|
| 260 |
-
return final
|
| 261 |
|
| 262 |
def run_and_submit_all( profile: gr.OAuthProfile | None):
|
| 263 |
"""
|
|
@@ -278,13 +241,12 @@ def run_and_submit_all( profile: gr.OAuthProfile | None):
|
|
| 278 |
questions_url = f"{api_url}/questions"
|
| 279 |
submit_url = f"{api_url}/submit"
|
| 280 |
|
| 281 |
-
# 1. Instantiate Agent
|
| 282 |
try:
|
| 283 |
agent = BasicAgent()
|
| 284 |
except Exception as e:
|
| 285 |
print(f"Error instantiating agent: {e}")
|
| 286 |
return f"Error initializing agent: {e}", None
|
| 287 |
-
# In the case of an app running as a hugging Face space, this link points toward your codebase ( usefull for others so please keep it public)
|
| 288 |
agent_code = f"https://huggingface.co/spaces/{space_id}/tree/main"
|
| 289 |
print(agent_code)
|
| 290 |
|
|
@@ -313,7 +275,7 @@ def run_and_submit_all( profile: gr.OAuthProfile | None):
|
|
| 313 |
results_log = []
|
| 314 |
answers_payload = []
|
| 315 |
print(f"Running agent on {len(questions_data)} questions...")
|
| 316 |
-
for
|
| 317 |
task_id = item.get("task_id")
|
| 318 |
question_text = item.get("question")
|
| 319 |
if not task_id or question_text is None:
|
|
@@ -401,7 +363,6 @@ with gr.Blocks() as demo:
|
|
| 401 |
run_button = gr.Button("Run Evaluation & Submit All Answers")
|
| 402 |
|
| 403 |
status_output = gr.Textbox(label="Run Status / Submission Result", lines=5, interactive=False)
|
| 404 |
-
# Removed max_rows=10 from DataFrame constructor
|
| 405 |
results_table = gr.DataFrame(label="Questions and Agent Answers", wrap=True)
|
| 406 |
|
| 407 |
run_button.click(
|
|
@@ -411,7 +372,6 @@ with gr.Blocks() as demo:
|
|
| 411 |
|
| 412 |
if __name__ == "__main__":
|
| 413 |
print("\n" + "-"*30 + " App Starting " + "-"*30)
|
| 414 |
-
# Check for SPACE_HOST and SPACE_ID at startup for information
|
| 415 |
space_host_startup = os.getenv("SPACE_HOST")
|
| 416 |
space_id_startup = os.getenv("SPACE_ID") # Get SPACE_ID at startup
|
| 417 |
|
|
|
|
| 3 |
import requests
|
| 4 |
import pandas as pd
|
| 5 |
import time
|
|
|
|
| 6 |
import re
|
| 7 |
import math
|
| 8 |
+
from typing import Dict, Any, List, Optional, Tuple
|
| 9 |
|
| 10 |
# (Keep Constants as is)
|
| 11 |
# --- Constants ---
|
| 12 |
DEFAULT_API_URL = "https://agents-course-unit4-scoring.hf.space"
|
| 13 |
|
| 14 |
+
# ------- Lightweight utilities -------
|
| 15 |
|
| 16 |
+
def maybe_compute_locally(question: str) -> Optional[str]:
|
| 17 |
"""
|
| 18 |
+
Fast, safe arithmetic for simple GAIA-style math.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 19 |
"""
|
| 20 |
+
text = question.strip().lower()
|
| 21 |
+
match = re.search(r"(?:calculate|compute|what is|evaluate)\s*[:\-]?\s*(.+)", text)
|
| 22 |
+
expr = match.group(1).strip() if match else text
|
| 23 |
+
expr = expr.replace("×", "*").replace("÷", "/").replace("^", "**")
|
| 24 |
+
expr = re.sub(r"(\d+)\s+percent", r"(\1/100)", expr)
|
| 25 |
+
# Only allow digits, ops, dots, spaces, parentheses
|
| 26 |
+
clean = expr.replace("**", "")
|
| 27 |
+
if not re.fullmatch(r"[0-9\.\s\+\-\*\/\(\)]+", clean):
|
| 28 |
+
return None
|
| 29 |
+
try:
|
| 30 |
+
result = eval(expr, {"__builtins__": {}}, {"sqrt": math.sqrt, "pow": pow})
|
| 31 |
+
except Exception:
|
| 32 |
+
return None
|
| 33 |
+
if isinstance(result, float):
|
| 34 |
+
result = 0.0 if abs(result) < 1e-12 else round(result, 6)
|
| 35 |
+
return str(result)
|
| 36 |
|
| 37 |
+
# ------- Wikipedia retrieval (fast, no API key) -------
|
| 38 |
+
|
| 39 |
+
WIKI_API = "https://en.wikipedia.org/w/api.php"
|
| 40 |
+
|
| 41 |
+
def wiki_search(query: str, limit: int = 3, timeout: int = 10) -> List[Dict[str, Any]]:
|
| 42 |
+
params = {
|
| 43 |
+
"action": "query",
|
| 44 |
+
"list": "search",
|
| 45 |
+
"srsearch": query,
|
| 46 |
+
"srlimit": limit,
|
| 47 |
+
"utf8": "1",
|
| 48 |
+
"format": "json",
|
| 49 |
+
}
|
| 50 |
+
r = requests.get(WIKI_API, params=params, timeout=timeout)
|
| 51 |
+
r.raise_for_status()
|
| 52 |
+
data = r.json()
|
| 53 |
+
return data.get("query", {}).get("search", [])
|
| 54 |
+
|
| 55 |
+
def wiki_extracts_by_pageids(pageids: List[int], timeout: int = 10) -> Dict[int, str]:
|
| 56 |
+
if not pageids:
|
| 57 |
+
return {}
|
| 58 |
+
params = {
|
| 59 |
+
"action": "query",
|
| 60 |
+
"prop": "extracts",
|
| 61 |
+
"explaintext": "1",
|
| 62 |
+
"exintro": "1",
|
| 63 |
+
"pageids": "|".join(str(pid) for pid in pageids),
|
| 64 |
+
"format": "json",
|
| 65 |
+
"utf8": "1",
|
| 66 |
+
}
|
| 67 |
+
r = requests.get(WIKI_API, params=params, timeout=timeout)
|
| 68 |
+
r.raise_for_status()
|
| 69 |
+
data = r.json()
|
| 70 |
+
pages = data.get("query", {}).get("pages", {})
|
| 71 |
+
out: Dict[int, str] = {}
|
| 72 |
+
for pid, meta in pages.items():
|
| 73 |
+
try:
|
| 74 |
+
out[int(pid)] = (meta.get("extract") or "").strip()
|
| 75 |
+
except Exception:
|
| 76 |
+
continue
|
| 77 |
+
return out
|
| 78 |
+
|
| 79 |
+
def retrieve_wikipedia_snippets(query: str, k: int = 3) -> List[Tuple[str, str]]:
|
| 80 |
+
"""
|
| 81 |
+
Returns list of (title, snippet_text).
|
| 82 |
+
"""
|
| 83 |
+
try:
|
| 84 |
+
hits = wiki_search(query, limit=k)
|
| 85 |
+
pageids = [h.get("pageid") for h in hits if isinstance(h.get("pageid"), int)]
|
| 86 |
+
extracts = wiki_extracts_by_pageids(pageids)
|
| 87 |
+
results: List[Tuple[str, str]] = []
|
| 88 |
+
for h in hits:
|
| 89 |
+
pid = h.get("pageid")
|
| 90 |
+
title = h.get("title", "")
|
| 91 |
+
extract = extracts.get(pid, "") if isinstance(pid, int) else ""
|
| 92 |
+
snippet = extract if extract else re.sub("<.*?>", "", h.get("snippet", "")).strip()
|
| 93 |
+
if title and snippet:
|
| 94 |
+
# take first 2 sentences to be concise
|
| 95 |
+
first_two = re.split(r"(?<=[.!?])\s+", snippet)[:2]
|
| 96 |
+
results.append((title, " ".join(first_two)))
|
| 97 |
+
return results
|
| 98 |
+
except Exception:
|
| 99 |
+
return []
|
| 100 |
+
|
| 101 |
+
# ------- Minimal LLM client (single provider, optional) -------
|
| 102 |
+
|
| 103 |
+
class LLMClient:
|
| 104 |
+
"""
|
| 105 |
+
Single-shot caller with tiny retry. Defaults to OpenRouter (cheapest path).
|
| 106 |
+
Set one of:
|
| 107 |
+
- OPENAI_API_KEY + OPENAI_MODEL + OPENAI_BASE_URL
|
| 108 |
+
- or OPENROUTER_API_KEY + OPENROUTER_MODEL (default used if key present)
|
| 109 |
+
Priority: OPENAI if key present, else OPENROUTER if key present, else disabled.
|
| 110 |
+
"""
|
| 111 |
def __init__(self):
|
| 112 |
+
self.provider: Optional[str] = None
|
| 113 |
+
self.headers: Dict[str, str] = {}
|
| 114 |
+
self.base_url: Optional[str] = None
|
| 115 |
+
self.model: Optional[str] = None
|
| 116 |
|
| 117 |
+
openai_key = os.getenv("OPENAI_API_KEY")
|
| 118 |
+
openrouter_key = os.getenv("OPENROUTER_API_KEY")
|
| 119 |
+
|
| 120 |
+
if openai_key:
|
| 121 |
+
self.provider = "openai"
|
|
|
|
| 122 |
self.base_url = os.getenv("OPENAI_BASE_URL", "https://api.openai.com/v1")
|
| 123 |
+
self.model = os.getenv("OPENAI_MODEL", "gpt-4o-mini")
|
| 124 |
+
self.headers = {"Authorization": f"Bearer {openai_key}", "Content-Type": "application/json"}
|
| 125 |
+
elif openrouter_key:
|
| 126 |
+
self.provider = "openrouter"
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 127 |
self.base_url = "https://openrouter.ai/api/v1"
|
| 128 |
+
self.model = os.getenv("OPENROUTER_MODEL", "meta-llama/llama-3.1-8b-instruct:free")
|
| 129 |
+
self.headers = {"Authorization": f"Bearer {openrouter_key}", "Content-Type": "application/json"}
|
|
|
|
|
|
|
|
|
|
| 130 |
else:
|
| 131 |
+
self.provider = None
|
| 132 |
|
| 133 |
+
def available(self) -> bool:
|
| 134 |
+
return self.provider is not None
|
|
|
|
|
|
|
| 135 |
|
| 136 |
+
def chat(self, system_prompt: str, user_prompt: str, max_tokens: int = 192, temperature: float = 0.3) -> str:
|
| 137 |
+
if not self.available():
|
| 138 |
+
return "LLM unavailable"
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 139 |
url = f"{self.base_url}/chat/completions"
|
| 140 |
+
payload = {
|
| 141 |
"model": self.model,
|
| 142 |
"messages": [
|
| 143 |
{"role": "system", "content": system_prompt},
|
|
|
|
| 145 |
],
|
| 146 |
"max_tokens": max_tokens,
|
| 147 |
"temperature": temperature,
|
| 148 |
+
"top_p": 0.9,
|
| 149 |
}
|
| 150 |
+
for attempt in range(3):
|
| 151 |
+
try:
|
| 152 |
+
resp = requests.post(url, headers=self.headers, json=payload, timeout=60)
|
| 153 |
+
if resp.status_code in (429, 500, 502, 503, 504):
|
| 154 |
+
time.sleep(1.5 * (attempt + 1))
|
| 155 |
+
continue
|
| 156 |
+
resp.raise_for_status()
|
| 157 |
+
data = resp.json()
|
| 158 |
+
text = data.get("choices", [{}])[0].get("message", {}).get("content", "")
|
| 159 |
+
return (text or "").strip()
|
| 160 |
+
except Exception as e:
|
| 161 |
+
if attempt == 2:
|
| 162 |
+
return f"LLM error: {e}"
|
| 163 |
+
time.sleep(1.0)
|
| 164 |
+
return "LLM error"
|
| 165 |
+
|
| 166 |
+
# --- Basic Agent Definition (lean) ---
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 167 |
|
| 168 |
class BasicAgent:
|
| 169 |
def __init__(self):
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 170 |
self.client = LLMClient()
|
| 171 |
+
self.max_new_tokens = int(os.getenv("MAX_NEW_TOKENS", "192"))
|
| 172 |
+
print(f"BasicAgent initialized. LLM={'on' if self.client.available() else 'off'}")
|
| 173 |
|
| 174 |
def _finalize(self, text: str) -> str:
|
| 175 |
+
t = text.strip().strip("`")
|
|
|
|
|
|
|
|
|
|
|
|
|
| 176 |
t = re.sub(r"\s+", " ", t)
|
| 177 |
+
return t
|
| 178 |
|
| 179 |
+
def __call__(self, question: str) -> str:
|
| 180 |
+
print(f"Agent received question (first 50 chars): {question[:50]}...")
|
| 181 |
+
|
| 182 |
+
# 1) Fast local math if applicable
|
| 183 |
local = maybe_compute_locally(question)
|
| 184 |
if local is not None:
|
| 185 |
+
print("Answered via local compute.")
|
| 186 |
return local
|
| 187 |
|
| 188 |
+
# 2) Wikipedia retrieve
|
| 189 |
+
snippets = retrieve_wikipedia_snippets(question, k=3)
|
| 190 |
+
|
| 191 |
+
# 3) If no LLM available, respond with best snippet sentence
|
| 192 |
+
if not self.client.available():
|
| 193 |
+
if snippets:
|
| 194 |
+
title, snip = snippets[0]
|
| 195 |
+
ans = self._finalize(snip)
|
| 196 |
+
print(f"Answered via Wikipedia snippet ({title}).")
|
| 197 |
+
return ans
|
| 198 |
+
print("No LLM and no snippets; returning fallback.")
|
| 199 |
+
return "I cannot determine the answer."
|
| 200 |
+
|
| 201 |
+
# 4) Compose minimal prompt with top snippets to keep tokens low
|
| 202 |
+
context_lines = []
|
| 203 |
+
for title, snip in snippets[:3]:
|
| 204 |
+
context_lines.append(f"- {title}: {snip}")
|
| 205 |
+
context = "\n".join(context_lines) if context_lines else "(no external context found)"
|
| 206 |
+
|
| 207 |
+
system_prompt = (
|
| 208 |
+
"You answer factually and concisely. Use the provided context if helpful. "
|
| 209 |
+
"Return only the final answer, no extra text."
|
| 210 |
)
|
| 211 |
+
user_prompt = f"Question: {question}\nContext:\n{context}\nAnswer:"
|
| 212 |
|
| 213 |
+
reply = self.client.chat(system_prompt, user_prompt, max_tokens=self.max_new_tokens, temperature=0.3)
|
| 214 |
+
answer = self._finalize(reply)
|
| 215 |
|
| 216 |
+
# 5) If the model failed, fall back to best snippet
|
| 217 |
+
if not answer or answer.lower().startswith("llm error"):
|
| 218 |
+
if snippets:
|
| 219 |
+
answer = self._finalize(snippets[0][1])
|
| 220 |
+
else:
|
| 221 |
+
answer = "I cannot determine the answer."
|
| 222 |
+
print(f"Agent returning answer: {answer[:80]}")
|
| 223 |
+
return answer
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 224 |
|
| 225 |
def run_and_submit_all( profile: gr.OAuthProfile | None):
|
| 226 |
"""
|
|
|
|
| 241 |
questions_url = f"{api_url}/questions"
|
| 242 |
submit_url = f"{api_url}/submit"
|
| 243 |
|
| 244 |
+
# 1. Instantiate Agent
|
| 245 |
try:
|
| 246 |
agent = BasicAgent()
|
| 247 |
except Exception as e:
|
| 248 |
print(f"Error instantiating agent: {e}")
|
| 249 |
return f"Error initializing agent: {e}", None
|
|
|
|
| 250 |
agent_code = f"https://huggingface.co/spaces/{space_id}/tree/main"
|
| 251 |
print(agent_code)
|
| 252 |
|
|
|
|
| 275 |
results_log = []
|
| 276 |
answers_payload = []
|
| 277 |
print(f"Running agent on {len(questions_data)} questions...")
|
| 278 |
+
for item in questions_data:
|
| 279 |
task_id = item.get("task_id")
|
| 280 |
question_text = item.get("question")
|
| 281 |
if not task_id or question_text is None:
|
|
|
|
| 363 |
run_button = gr.Button("Run Evaluation & Submit All Answers")
|
| 364 |
|
| 365 |
status_output = gr.Textbox(label="Run Status / Submission Result", lines=5, interactive=False)
|
|
|
|
| 366 |
results_table = gr.DataFrame(label="Questions and Agent Answers", wrap=True)
|
| 367 |
|
| 368 |
run_button.click(
|
|
|
|
| 372 |
|
| 373 |
if __name__ == "__main__":
|
| 374 |
print("\n" + "-"*30 + " App Starting " + "-"*30)
|
|
|
|
| 375 |
space_host_startup = os.getenv("SPACE_HOST")
|
| 376 |
space_id_startup = os.getenv("SPACE_ID") # Get SPACE_ID at startup
|
| 377 |
|