OhBrian's picture
更新为 LangGraph 工作流模式
d6b82a1
Raw
History Blame Contribute Delete
8.36 kB
import os
import gradio as gr
import pandas as pd
import requests
from agent import GaiaAgent
from config import (
DEFAULT_API_URL,
HF_PLANNER_USE_RESPONSE_FORMAT,
HF_TEXT_MODEL,
HF_VISION_MODEL,
)
def run_and_submit_all(profile: gr.OAuthProfile | None):
"""
拉取全部问题,使用 GaiaAgent 逐题作答,提交全部答案,并展示评测结果。
"""
space_id = os.getenv("SPACE_ID")
if profile:
username = f"{profile.username}"
print(f"用户已登录:{username}")
else:
print("用户未登录。")
return "请先点击按钮登录 Hugging Face。", None
api_url = DEFAULT_API_URL
questions_url = f"{api_url}/questions"
submit_url = f"{api_url}/submit"
try:
agent = GaiaAgent()
except Exception as e:
print(f"实例化 Agent 时出错:{e}")
return f"Agent 初始化失败:{e}", None
agent_code = (
f"https://huggingface.co/spaces/{space_id}/tree/main"
if space_id
else "本地运行,未提供 SPACE_ID。"
)
print(agent_code)
print(f"正在从以下地址拉取问题:{questions_url}")
try:
response = requests.get(questions_url, timeout=15)
response.raise_for_status()
questions_data = response.json()
if not questions_data:
print("拉取到的问题列表为空。")
return "拉取到的问题列表为空,或返回格式无效。", None
print(f"已拉取 {len(questions_data)} 个问题。")
except requests.exceptions.RequestException as e:
print(f"拉取问题时发生网络错误:{e}")
return f"拉取问题失败:{e}", None
except requests.exceptions.JSONDecodeError as e:
print(f"解析问题接口的 JSON 响应时出错:{e}")
print(f"响应文本:{response.text[:500]}")
return f"解析问题接口响应失败:{e}", None
except Exception as e:
print(f"拉取问题时发生未预期错误:{e}")
return f"拉取问题时发生未预期错误:{e}", None
results_log = []
answers_payload = []
print(f"正在让 Agent 处理 {len(questions_data)} 个问题...")
for item in questions_data:
task_id = item.get("task_id")
question_text = item.get("question")
file_name = item.get("file_name") or ""
if not task_id or question_text is None:
print(f"跳过缺少 task_id 或 question 的条目:{item}")
continue
try:
result = agent.answer_task(question_text, task_id=task_id, file_name=file_name)
submitted_answer = result.answer or "无法确定"
answers_payload.append(
{"task_id": task_id, "submitted_answer": submitted_answer}
)
results_log.append(
{
"任务 ID": task_id,
"附件": file_name,
"问题": question_text,
"提交答案": submitted_answer,
"来源": result.source,
"置信度": result.confidence,
"错误": result.error,
}
)
except Exception as e:
print(f"Agent 处理任务 {task_id} 时出错:{e}")
results_log.append(
{
"任务 ID": task_id,
"附件": file_name,
"问题": question_text,
"提交答案": f"AGENT ERROR: {e}",
"来源": "agent.exception",
"置信度": "low",
"错误": str(e),
}
)
if not answers_payload:
print("Agent 没有生成任何可提交的答案。")
return "Agent 没有生成任何可提交的答案。", pd.DataFrame(results_log)
submission_data = {
"username": username.strip(),
"agent_code": agent_code,
"answers": answers_payload,
}
status_update = f"Agent 已完成作答。正在为用户 '{username}' 提交 {len(answers_payload)} 个答案..."
print(status_update)
print(f"正在向以下地址提交 {len(answers_payload)} 个答案:{submit_url}")
try:
response = requests.post(submit_url, json=submission_data, timeout=60)
response.raise_for_status()
result_data = response.json()
final_status = (
f"提交成功!\n"
f"用户:{result_data.get('username')}\n"
f"总分:{result_data.get('score', 'N/A')}% "
f"({result_data.get('correct_count', '?')}/{result_data.get('total_attempted', '?')} 正确)\n"
f"消息:{result_data.get('message', '未收到消息。')}"
)
print("提交成功。")
return final_status, pd.DataFrame(results_log)
except requests.exceptions.HTTPError as e:
error_detail = f"服务器返回状态码 {e.response.status_code}。"
try:
error_json = e.response.json()
error_detail += f" 详情:{error_json.get('detail', e.response.text)}"
except requests.exceptions.JSONDecodeError:
error_detail += f" 响应:{e.response.text[:500]}"
status_message = f"提交失败:{error_detail}"
print(status_message)
return status_message, pd.DataFrame(results_log)
except requests.exceptions.Timeout:
status_message = "提交失败:请求超时。"
print(status_message)
return status_message, pd.DataFrame(results_log)
except requests.exceptions.RequestException as e:
status_message = f"提交失败:网络错误 - {e}"
print(status_message)
return status_message, pd.DataFrame(results_log)
except Exception as e:
status_message = f"提交过程中发生未预期错误:{e}"
print(status_message)
return status_message, pd.DataFrame(results_log)
with gr.Blocks() as demo:
gr.Markdown("# LangGraph GAIA Agent 评测运行器")
gr.Markdown(
f"""
**当前策略:**
1. LangGraph 先调用 LLM 做题型分类,明确应该使用哪类工具。
2. 工作流按题型进入规则、Python、Excel、Wikipedia、体育统计、网页读取/搜索等处理节点。
3. 工具结果统一交给 LLM 做最终答案格式化;LLM 失败时使用工具候选答案兜底。
4. 未识别类型或工具无结果时进入兜底节点,按低成本工具顺序尝试。
5. 音频、视频和棋局图工具当前禁用;这类题会返回“无法确定”。
6. 当前文本模型:`{HF_TEXT_MODEL}`;当前视觉模型:`{HF_VISION_MODEL or "未启用"}`。
7. 强制 JSON response_format:`{HF_PLANNER_USE_RESPONSE_FORMAT}`。
**提交前检查:**
Space Secrets 至少需要 `HF_TOKEN`,且该 token 必须允许调用 Inference Providers。
"""
)
gr.LoginButton()
run_button = gr.Button("运行评测并提交全部答案")
status_output = gr.Textbox(label="运行状态 / 提交结果", lines=5, interactive=False)
results_table = gr.DataFrame(label="问题与 Agent 答案", wrap=True)
run_button.click(fn=run_and_submit_all, outputs=[status_output, results_table])
if __name__ == "__main__":
print("\n" + "-" * 30 + " 应用启动中 " + "-" * 30)
space_host_startup = os.getenv("SPACE_HOST")
space_id_startup = os.getenv("SPACE_ID")
if space_host_startup:
print(f"找到 SPACE_HOST:{space_host_startup}")
print(f" 运行地址应为:https://{space_host_startup}.hf.space")
else:
print("未找到 SPACE_HOST 环境变量(可能是在本地运行)。")
if space_id_startup:
print(f"找到 SPACE_ID:{space_id_startup}")
print(f" 仓库地址:https://huggingface.co/spaces/{space_id_startup}")
print(f" 代码树地址:https://huggingface.co/spaces/{space_id_startup}/tree/main")
else:
print("未找到 SPACE_ID 环境变量(可能是在本地运行)。无法确定仓库地址。")
print("-" * (60 + len(" 应用启动中 ")) + "\n")
print("正在启动 LangGraph GAIA Agent 评测 Gradio 界面...")
demo.launch(debug=True, share=False)