核心修复: - knowledge/base.py: BM25Index.add_documents 从覆盖改为追加+去重, 修复只有最后上传文件的 chunks 保留在 BM25 中的严重 bug (影响: 2.docx/3.docx/PDF 的 755 个 chunk 在 BM25 中完全缺失) 检索增强 (延续上次会话): - core/engine.py: section cluster boost + lexical match exemption - api/chat_routes.py: lexical/cluster rescue 层 + SSE 事件 - core/mmr.py: MMR 去重改进 评测体系: - tests/eval_dataset_v2.json: 62 题综合评测集 (9 种题型×4 文档) - scripts/eval_e2e.py: 推理模型 LLM 评分兼容 + 新数据集格式支持 - scripts/validate_eval_dataset.py: 数据集验证工具 其他: - parsers/mineru_parser.py: 解析器改进
742 lines
25 KiB
Python
742 lines
25 KiB
Python
#!/usr/bin/env python
|
||
# -*- coding: utf-8 -*-
|
||
"""
|
||
RAG 端到端评测脚本 (eval_e2e.py)
|
||
|
||
通过 HTTP API 调用 /rag 接口,解析 SSE 流式响应,
|
||
对每个问题从关键词覆盖、LLM 质量评分两个维度进行评测,
|
||
支持基线录制与阶段间对比。
|
||
|
||
用法:
|
||
# 录制基线
|
||
python scripts/eval_e2e.py --baseline
|
||
|
||
# 阶段评测并与基线对比
|
||
python scripts/eval_e2e.py --phase 1
|
||
|
||
# 仅评测不对比
|
||
python scripts/eval_e2e.py --phase test
|
||
|
||
# 指定数据集和输出路径
|
||
python scripts/eval_e2e.py --dataset data/eval/eval_dataset.json --phase 0 --output data/eval_results/phase0.json
|
||
|
||
# 跳过 LLM 评分(快速模式)
|
||
python scripts/eval_e2e.py --phase test --no_llm
|
||
"""
|
||
|
||
import argparse
|
||
import json
|
||
import logging
|
||
import os
|
||
import re
|
||
import sys
|
||
import time
|
||
from datetime import datetime
|
||
from pathlib import Path
|
||
|
||
# 添加项目根目录到路径
|
||
PROJECT_ROOT = Path(__file__).parent.parent
|
||
sys.path.insert(0, str(PROJECT_ROOT))
|
||
|
||
logging.basicConfig(
|
||
level=logging.INFO,
|
||
format='%(asctime)s [%(levelname)s] %(message)s'
|
||
)
|
||
logger = logging.getLogger(__name__)
|
||
|
||
|
||
def setup_file_logging(log_path: str):
|
||
"""将日志同时输出到文件,避免 Windows 控制台编码问题"""
|
||
fh = logging.FileHandler(log_path, encoding='utf-8')
|
||
fh.setLevel(logging.INFO)
|
||
fh.setFormatter(logging.Formatter('%(asctime)s [%(levelname)s] %(message)s'))
|
||
logger.addHandler(fh)
|
||
|
||
# ───────────── 配置 ─────────────
|
||
RAG_API_URL = "http://127.0.0.1:5001/rag"
|
||
DEFAULT_DATASET = "tests/eval_dataset_v2.json"
|
||
RESULTS_DIR = "data/eval_results"
|
||
REQUEST_TIMEOUT = 120 # 秒
|
||
REQUEST_INTERVAL = 1.5 # 请求间隔(秒),避免过载
|
||
|
||
|
||
# ───────────── SSE 解析 ─────────────
|
||
def call_rag_sse(question: str, collections: list = None, api_url: str = None) -> dict:
|
||
"""
|
||
调用 /rag 接口,解析 SSE 流式响应,提取 finish 事件中的完整回答。
|
||
参考 docs/curl测试手册.md 中的调用规范。
|
||
|
||
Returns:
|
||
{
|
||
"answer": str,
|
||
"sources": list,
|
||
"citations": list,
|
||
"duration_ms": int,
|
||
"error": str | None
|
||
}
|
||
"""
|
||
import requests
|
||
|
||
headers = {
|
||
"Content-Type": "application/json",
|
||
"Accept": "text/event-stream",
|
||
"Authorization": "Bearer mock-token-admin"
|
||
}
|
||
# chat_history 在生产模式下是必填字段
|
||
payload = {"message": question, "chat_history": []}
|
||
if collections:
|
||
payload["collections"] = collections
|
||
|
||
result = {
|
||
"answer": "",
|
||
"sources": [],
|
||
"citations": [],
|
||
"duration_ms": 0,
|
||
"error": None
|
||
}
|
||
|
||
try:
|
||
resp = requests.post(
|
||
api_url or RAG_API_URL, json=payload, headers=headers,
|
||
stream=True, timeout=REQUEST_TIMEOUT
|
||
)
|
||
resp.raise_for_status()
|
||
|
||
# 逐行解析 SSE 事件
|
||
for raw_line in resp.iter_lines(decode_unicode=True):
|
||
if not raw_line or not raw_line.startswith("data:"):
|
||
continue
|
||
data_str = raw_line[len("data:"):].strip()
|
||
if not data_str:
|
||
continue
|
||
try:
|
||
event = json.loads(data_str)
|
||
except json.JSONDecodeError:
|
||
continue
|
||
|
||
evt_type = event.get("type")
|
||
if evt_type == "finish":
|
||
result["answer"] = event.get("answer", "")
|
||
result["sources"] = event.get("sources", [])
|
||
result["citations"] = event.get("citations", [])
|
||
result["duration_ms"] = event.get("duration_ms", 0)
|
||
break
|
||
elif evt_type == "error":
|
||
result["error"] = event.get("message", "未知错误")
|
||
break
|
||
|
||
except requests.exceptions.ConnectionError:
|
||
result["error"] = "无法连接到 RAG 服务,请确认服务已启动 (localhost:5001)"
|
||
except requests.exceptions.Timeout:
|
||
result["error"] = f"请求超时 ({REQUEST_TIMEOUT}s)"
|
||
except Exception as e:
|
||
result["error"] = str(e)
|
||
|
||
return result
|
||
|
||
|
||
# ───────────── 评分模块 ─────────────
|
||
def keyword_coverage(answer: str, expected_keywords: list) -> dict:
|
||
"""
|
||
关键词覆盖率评分。
|
||
检查 answer 中包含 expected_keywords 的百分比。
|
||
|
||
Returns:
|
||
{"score": 0~1, "matched": [...], "missing": [...]}
|
||
"""
|
||
if not expected_keywords or not answer:
|
||
return {"score": 0.0, "matched": [], "missing": expected_keywords or []}
|
||
|
||
matched = [kw for kw in expected_keywords if kw in answer]
|
||
missing = [kw for kw in expected_keywords if kw not in answer]
|
||
score = len(matched) / len(expected_keywords)
|
||
return {"score": round(score, 4), "matched": matched, "missing": missing}
|
||
|
||
|
||
def llm_quality_score(query: str, answer: str, reference: str) -> dict:
|
||
"""
|
||
使用 LLM 对回答质量进行多维度评分 (0-10)。
|
||
|
||
Returns:
|
||
{"accuracy": X, "completeness": X, "relevance": X, "fluency": X, "overall": X}
|
||
"""
|
||
try:
|
||
from openai import OpenAI
|
||
from config import DASHSCOPE_API_KEY, DASHSCOPE_BASE_URL, DASHSCOPE_MODEL
|
||
except ImportError:
|
||
logger.warning("无法导入 openai 或 config 模块,跳过 LLM 评分")
|
||
return {"overall": 0.0, "error": "import_failed"}
|
||
|
||
client = OpenAI(api_key=DASHSCOPE_API_KEY, base_url=DASHSCOPE_BASE_URL)
|
||
|
||
prompt = f"""你是一个专业的 RAG 系统评测员。请对以下回答进行质量评分。
|
||
|
||
【用户问题】
|
||
{query}
|
||
|
||
【参考答案(人工编写的高质量答案)】
|
||
{reference}
|
||
|
||
【系统实际回答】
|
||
{answer}
|
||
|
||
【评分维度】(每项 0-10 分)
|
||
1. 准确性(accuracy):系统回答中的事实信息是否与参考答案一致,有无错误信息
|
||
2. 完整性(completeness):是否覆盖了参考答案中的关键要点
|
||
3. 相关性(relevance):回答是否直接针对问题,没有跑题或冗余
|
||
4. 流畅性(fluency):回答是否通顺、结构清晰、易于理解
|
||
|
||
【输出要求】
|
||
只输出一个纯 JSON 对象,不要包含任何其他文字、解释或 markdown 格式:
|
||
{{"accuracy": <整数>, "completeness": <整数>, "relevance": <整数>, "fluency": <整数>, "overall": <整数>}}"""
|
||
|
||
def _extract_json(text: str) -> dict | None:
|
||
"""从文本中提取 JSON,支持嵌套大括号"""
|
||
if not text:
|
||
return None
|
||
# 1. 移除推理模型的思考标签及其内容
|
||
text = re.sub(r'<think>[\s\S]*?</think>', '', text).strip()
|
||
# 2. 尝试直接解析整个文本
|
||
try:
|
||
return json.loads(text)
|
||
except json.JSONDecodeError:
|
||
pass
|
||
# 3. 贪婪匹配最大的 {...} 块(支持嵌套)
|
||
# 从最后一个 } 往前找匹配的 {
|
||
brace_depth = 0
|
||
start = -1
|
||
end = -1
|
||
for i in range(len(text) - 1, -1, -1):
|
||
if text[i] == '}':
|
||
if brace_depth == 0:
|
||
end = i
|
||
brace_depth += 1
|
||
elif text[i] == '{':
|
||
brace_depth -= 1
|
||
if brace_depth == 0:
|
||
start = i
|
||
break
|
||
if start >= 0 and end > start:
|
||
try:
|
||
return json.loads(text[start:end + 1])
|
||
except json.JSONDecodeError:
|
||
pass
|
||
# 4. 回退:简单单层 {...} 匹配
|
||
m = re.search(r'\{[^{}]+\}', text)
|
||
if m:
|
||
try:
|
||
return json.loads(m.group())
|
||
except json.JSONDecodeError:
|
||
pass
|
||
return None
|
||
|
||
try:
|
||
# 构建请求参数
|
||
request_params = {
|
||
"model": DASHSCOPE_MODEL,
|
||
"messages": [
|
||
{"role": "system", "content": "You are a JSON-only evaluator. Output ONLY a valid JSON object, nothing else."},
|
||
{"role": "user", "content": prompt}
|
||
],
|
||
"temperature": 0.1,
|
||
"max_tokens": 2000,
|
||
}
|
||
# 尝试关闭推理模型的思考输出(部分 API 支持)
|
||
try:
|
||
request_params["extra_body"] = {"enable_thinking": False}
|
||
except Exception:
|
||
pass
|
||
|
||
response = client.chat.completions.create(**request_params)
|
||
msg = response.choices[0].message
|
||
|
||
# 优先取 content,如果为空则尝试 reasoning_content 后的内容
|
||
text = getattr(msg, 'content', '') or ''
|
||
|
||
# 如果 content 为空,尝试从 reasoning_content 中提取
|
||
# (部分推理模型 API 将思考内容放在 reasoning_content,回答放在 content)
|
||
if not text.strip():
|
||
reasoning = getattr(msg, 'reasoning_content', '') or ''
|
||
if reasoning:
|
||
# 从思考内容末尾尝试提取 JSON
|
||
text = reasoning
|
||
|
||
if not text.strip():
|
||
logger.warning("LLM 返回内容为空")
|
||
return {"overall": 0.0, "error": "empty_response"}
|
||
|
||
scores = _extract_json(text)
|
||
if scores:
|
||
# 归一化到 0-1
|
||
for k in list(scores.keys()):
|
||
try:
|
||
scores[k] = round(min(10, max(0, float(scores[k]))) / 10.0, 4)
|
||
except (ValueError, TypeError):
|
||
scores[k] = 0.0
|
||
return scores
|
||
else:
|
||
# 记录更多内容用于调试
|
||
debug_text = text[:200].replace('\n', '\\n')
|
||
logger.warning(f"LLM 返回内容无法解析 JSON: {debug_text}")
|
||
return {"overall": 0.0, "error": "parse_failed"}
|
||
except Exception as e:
|
||
logger.warning(f"LLM 评分异常: {e}")
|
||
return {"overall": 0.0, "error": str(e)}
|
||
|
||
|
||
# ───────────── 评测主流程 ─────────────
|
||
def evaluate_dataset(dataset_path: str, use_llm: bool = True, api_url: str = None) -> dict:
|
||
"""
|
||
对数据集中的所有问题进行评测。
|
||
|
||
Returns:
|
||
{
|
||
"meta": {...},
|
||
"questions": [ {id, query, answer, scores, ...}, ... ],
|
||
"summary": { avg_keyword, avg_llm, by_type, by_difficulty }
|
||
}
|
||
"""
|
||
with open(dataset_path, 'r', encoding='utf-8') as f:
|
||
dataset = json.load(f)
|
||
|
||
questions = dataset.get("queries", dataset.get("questions", []))
|
||
total = len(questions)
|
||
if total == 0:
|
||
logger.error("数据集中没有问题")
|
||
return {}
|
||
|
||
logger.info(f"开始评测,共 {total} 个问题...")
|
||
|
||
results = []
|
||
kw_scores = []
|
||
llm_scores = []
|
||
by_type = {}
|
||
by_difficulty = {}
|
||
|
||
for i, q in enumerate(questions, 1):
|
||
qid = q["id"]
|
||
query = q["query"]
|
||
qtype = q.get("query_type", "unknown")
|
||
difficulty = q.get("difficulty", "medium")
|
||
expected_keywords = q.get("expected_keywords", [])
|
||
reference_answer = q.get("reference_answer", "")
|
||
|
||
logger.info(f"[{i}/{total}] {qid}: {query[:40]}...")
|
||
|
||
# 调用 RAG
|
||
t0 = time.time()
|
||
rag_result = call_rag_sse(query, api_url=api_url)
|
||
elapsed = round(time.time() - t0, 2)
|
||
answer = rag_result.get("answer", "")
|
||
error = rag_result.get("error")
|
||
|
||
if error:
|
||
logger.warning(f" !! API 错误: {error}")
|
||
results.append({
|
||
"id": qid, "query": query, "query_type": qtype,
|
||
"difficulty": difficulty, "answer": "", "error": error,
|
||
"elapsed_s": elapsed,
|
||
"keyword_score": 0.0, "llm_scores": {"overall": 0.0}
|
||
})
|
||
continue
|
||
|
||
# 关键词覆盖评分
|
||
kw_result = keyword_coverage(answer, expected_keywords)
|
||
kw_score = kw_result["score"]
|
||
kw_scores.append(kw_score)
|
||
|
||
# LLM 质量评分
|
||
llm_result = {}
|
||
if use_llm and reference_answer:
|
||
llm_result = llm_quality_score(query, answer, reference_answer)
|
||
llm_scores.append(llm_result.get("overall", 0.0))
|
||
|
||
entry = {
|
||
"id": qid,
|
||
"query": query,
|
||
"query_type": qtype,
|
||
"difficulty": difficulty,
|
||
"answer": answer,
|
||
"sources": rag_result.get("sources", []),
|
||
"citations": rag_result.get("citations", []),
|
||
"duration_ms": rag_result.get("duration_ms", 0),
|
||
"elapsed_s": elapsed,
|
||
"keyword_score": kw_score,
|
||
"keyword_matched": kw_result["matched"],
|
||
"keyword_missing": kw_result["missing"],
|
||
"llm_scores": llm_result,
|
||
"reference_answer": reference_answer[:300]
|
||
}
|
||
results.append(entry)
|
||
|
||
# 按类型/难度分组
|
||
for group_dict, key in [(by_type, qtype), (by_difficulty, difficulty)]:
|
||
if key not in group_dict:
|
||
group_dict[key] = {"kw": [], "llm": []}
|
||
group_dict[key]["kw"].append(kw_score)
|
||
if "overall" in llm_result:
|
||
group_dict[key]["llm"].append(llm_result["overall"])
|
||
|
||
# 打印进度
|
||
llm_str = f", LLM={llm_result.get('overall', 'N/A')}" if use_llm else ""
|
||
logger.info(f" 关键词={kw_score:.0%}{llm_str}, 耗时={elapsed}s")
|
||
|
||
# 请求间隔
|
||
if i < total:
|
||
time.sleep(REQUEST_INTERVAL)
|
||
|
||
# 汇总
|
||
avg_kw = sum(kw_scores) / len(kw_scores) if kw_scores else 0.0
|
||
avg_llm = sum(llm_scores) / len(llm_scores) if llm_scores else 0.0
|
||
|
||
summary = {
|
||
"avg_keyword_coverage": round(avg_kw, 4),
|
||
"avg_llm_overall": round(avg_llm, 4),
|
||
"total_questions": total,
|
||
"successful": len(kw_scores),
|
||
"by_type": {
|
||
k: {
|
||
"avg_keyword": round(sum(v["kw"]) / len(v["kw"]), 4) if v["kw"] else 0,
|
||
"avg_llm": round(sum(v["llm"]) / len(v["llm"]), 4) if v["llm"] else 0,
|
||
"count": len(v["kw"])
|
||
}
|
||
for k, v in by_type.items()
|
||
},
|
||
"by_difficulty": {
|
||
k: {
|
||
"avg_keyword": round(sum(v["kw"]) / len(v["kw"]), 4) if v["kw"] else 0,
|
||
"avg_llm": round(sum(v["llm"]) / len(v["llm"]), 4) if v["llm"] else 0,
|
||
"count": len(v["kw"])
|
||
}
|
||
for k, v in by_difficulty.items()
|
||
}
|
||
}
|
||
|
||
return {
|
||
"meta": {
|
||
"timestamp": datetime.now().isoformat(),
|
||
"dataset": dataset_path,
|
||
"use_llm": use_llm
|
||
},
|
||
"questions": results,
|
||
"summary": summary
|
||
}
|
||
|
||
|
||
# ───────────── 对比模块 ─────────────
|
||
def compare_results(baseline_path: str, current_path: str) -> dict:
|
||
"""
|
||
对比基线和当前阶段的评测结果。
|
||
|
||
Returns:
|
||
{
|
||
"per_question": [{id, query, baseline_kw, current_kw, delta_kw, ...}],
|
||
"summary_delta": { keyword_delta, llm_delta },
|
||
"regressions": [{id, query, delta_kw, delta_llm}]
|
||
}
|
||
"""
|
||
with open(baseline_path, 'r', encoding='utf-8') as f:
|
||
baseline = json.load(f)
|
||
with open(current_path, 'r', encoding='utf-8') as f:
|
||
current = json.load(f)
|
||
|
||
# 按 ID 索引
|
||
b_map = {q["id"]: q for q in baseline.get("questions", [])}
|
||
c_map = {q["id"]: q for q in current.get("questions", [])}
|
||
|
||
per_question = []
|
||
regressions = []
|
||
|
||
for qid in sorted(b_map.keys()):
|
||
bq = b_map[qid]
|
||
cq = c_map.get(qid)
|
||
if not cq:
|
||
continue
|
||
|
||
b_kw = bq.get("keyword_score", 0)
|
||
c_kw = cq.get("keyword_score", 0)
|
||
b_llm = bq.get("llm_scores", {}).get("overall", 0)
|
||
c_llm = cq.get("llm_scores", {}).get("overall", 0)
|
||
|
||
delta_kw = round(c_kw - b_kw, 4)
|
||
delta_llm = round(c_llm - b_llm, 4)
|
||
|
||
entry = {
|
||
"id": qid,
|
||
"query": bq.get("query", ""),
|
||
"query_type": bq.get("query_type", ""),
|
||
"difficulty": bq.get("difficulty", ""),
|
||
"baseline_kw": b_kw,
|
||
"current_kw": c_kw,
|
||
"delta_kw": delta_kw,
|
||
"baseline_llm": b_llm,
|
||
"current_llm": c_llm,
|
||
"delta_llm": delta_llm
|
||
}
|
||
per_question.append(entry)
|
||
|
||
# 回归判定:关键词覆盖或 LLM 评分下降超过 10%
|
||
if delta_kw < -0.1 or delta_llm < -0.1:
|
||
regressions.append(entry)
|
||
|
||
bs = baseline.get("summary", {})
|
||
cs = current.get("summary", {})
|
||
summary_delta = {
|
||
"keyword_delta": round(
|
||
cs.get("avg_keyword_coverage", 0) - bs.get("avg_keyword_coverage", 0), 4
|
||
),
|
||
"llm_delta": round(
|
||
cs.get("avg_llm_overall", 0) - bs.get("avg_llm_overall", 0), 4
|
||
),
|
||
"baseline_avg_kw": bs.get("avg_keyword_coverage", 0),
|
||
"current_avg_kw": cs.get("avg_keyword_coverage", 0),
|
||
"baseline_avg_llm": bs.get("avg_llm_overall", 0),
|
||
"current_avg_llm": cs.get("avg_llm_overall", 0)
|
||
}
|
||
|
||
return {
|
||
"per_question": per_question,
|
||
"summary_delta": summary_delta,
|
||
"regressions": regressions
|
||
}
|
||
|
||
|
||
# ───────────── 报告生成 ─────────────
|
||
def generate_report(eval_result: dict, comparison: dict = None, phase_name: str = "") -> str:
|
||
"""生成 Markdown 格式的评测报告"""
|
||
lines = []
|
||
meta = eval_result.get("meta", {})
|
||
summary = eval_result.get("summary", {})
|
||
questions = eval_result.get("questions", [])
|
||
|
||
lines.append(f"## RAG 端到端评测报告 — {phase_name or '未命名'}")
|
||
lines.append("")
|
||
lines.append(f"**评测时间**: {meta.get('timestamp', 'N/A')}")
|
||
lines.append(f"**数据集**: {meta.get('dataset', 'N/A')}")
|
||
lines.append(f"**LLM 评分**: {'启用' if meta.get('use_llm') else '跳过'}")
|
||
lines.append("")
|
||
|
||
# 汇总
|
||
lines.append("### 汇总指标")
|
||
lines.append("")
|
||
lines.append(f"| 指标 | 值 |")
|
||
lines.append(f"|------|-----|")
|
||
lines.append(f"| 平均关键词覆盖率 | {summary.get('avg_keyword_coverage', 0):.2%} |")
|
||
lines.append(f"| 平均 LLM 总分 | {summary.get('avg_llm_overall', 0):.2f} |")
|
||
lines.append(f"| 总问题数 | {summary.get('total_questions', 0)} |")
|
||
lines.append(f"| 成功评测数 | {summary.get('successful', 0)} |")
|
||
lines.append("")
|
||
|
||
# 按类型
|
||
by_type = summary.get("by_type", {})
|
||
if by_type:
|
||
lines.append("### 按查询类型")
|
||
lines.append("")
|
||
lines.append("| 类型 | 数量 | 平均关键词 | 平均 LLM |")
|
||
lines.append("|------|------|-----------|---------|")
|
||
for t, v in sorted(by_type.items()):
|
||
lines.append(f"| {t} | {v['count']} | {v['avg_keyword']:.2%} | {v['avg_llm']:.2f} |")
|
||
lines.append("")
|
||
|
||
# 按难度
|
||
by_diff = summary.get("by_difficulty", {})
|
||
if by_diff:
|
||
lines.append("### 按难度")
|
||
lines.append("")
|
||
lines.append("| 难度 | 数量 | 平均关键词 | 平均 LLM |")
|
||
lines.append("|------|------|-----------|---------|")
|
||
for d, v in sorted(by_diff.items()):
|
||
lines.append(f"| {d} | {v['count']} | {v['avg_keyword']:.2%} | {v['avg_llm']:.2f} |")
|
||
lines.append("")
|
||
|
||
# 对比结果
|
||
if comparison:
|
||
sd = comparison.get("summary_delta", {})
|
||
lines.append("### 与基线对比")
|
||
lines.append("")
|
||
lines.append(f"| 指标 | 基线 | 当前 | 变化 |")
|
||
lines.append(f"|------|------|------|------|")
|
||
lines.append(
|
||
f"| 关键词覆盖率 | {sd.get('baseline_avg_kw', 0):.2%} "
|
||
f"| {sd.get('current_avg_kw', 0):.2%} "
|
||
f"| {sd.get('keyword_delta', 0):+.2%} |"
|
||
)
|
||
lines.append(
|
||
f"| LLM 总分 | {sd.get('baseline_avg_llm', 0):.2f} "
|
||
f"| {sd.get('current_avg_llm', 0):.2f} "
|
||
f"| {sd.get('llm_delta', 0):+.2f} |"
|
||
)
|
||
lines.append("")
|
||
|
||
regressions = comparison.get("regressions", [])
|
||
if regressions:
|
||
lines.append(f"### 回归问题({len(regressions)} 个)")
|
||
lines.append("")
|
||
for r in regressions:
|
||
lines.append(
|
||
f"- **{r['id']}** {r['query'][:50]}... "
|
||
f"(关键词 {r['delta_kw']:+.0%}, LLM {r['delta_llm']:+.2f})"
|
||
)
|
||
lines.append("")
|
||
else:
|
||
lines.append("### 无回归问题")
|
||
lines.append("")
|
||
|
||
# 逐题详情
|
||
lines.append("### 逐题结果")
|
||
lines.append("")
|
||
for q in questions:
|
||
status = "ERROR" if q.get("error") else "OK"
|
||
llm_overall = q.get("llm_scores", {}).get("overall", "N/A")
|
||
if isinstance(llm_overall, float):
|
||
llm_overall = f"{llm_overall:.2f}"
|
||
lines.append(f"**{q['id']}** [{status}] {q['query']}")
|
||
lines.append(f"- 关键词覆盖: {q.get('keyword_score', 0):.0%}")
|
||
lines.append(f"- LLM 总分: {llm_overall}")
|
||
if q.get("keyword_missing"):
|
||
lines.append(f"- 缺失关键词: {', '.join(q['keyword_missing'])}")
|
||
if q.get("error"):
|
||
lines.append(f"- 错误: {q['error']}")
|
||
# 截取回答前 150 字
|
||
ans = q.get("answer", "")
|
||
if ans:
|
||
lines.append(f"- 回答摘要: {ans[:150]}...")
|
||
lines.append("")
|
||
|
||
return "\n".join(lines)
|
||
|
||
|
||
# ───────────── 主入口 ─────────────
|
||
def main():
|
||
parser = argparse.ArgumentParser(description="RAG 端到端评测脚本")
|
||
parser.add_argument(
|
||
"--dataset", type=str, default=DEFAULT_DATASET,
|
||
help=f"评测数据集路径 (默认: {DEFAULT_DATASET})"
|
||
)
|
||
parser.add_argument(
|
||
"--phase", type=str, default=None,
|
||
help="当前阶段名称,如 baseline / 0 / 1 / 2 ..."
|
||
)
|
||
parser.add_argument(
|
||
"--baseline", action="store_true",
|
||
help="录制基线(等价于 --phase baseline)"
|
||
)
|
||
parser.add_argument(
|
||
"--output", type=str, default=None,
|
||
help="结果 JSON 输出路径(默认自动生成)"
|
||
)
|
||
parser.add_argument(
|
||
"--no_llm", action="store_true",
|
||
help="跳过 LLM 评分(快速模式)"
|
||
)
|
||
parser.add_argument(
|
||
"--api_url", type=str, default=None,
|
||
help=f"RAG API 地址 (默认: {RAG_API_URL})"
|
||
)
|
||
args = parser.parse_args()
|
||
|
||
# Windows 编码
|
||
if sys.platform == 'win32':
|
||
sys.stdout.reconfigure(encoding='utf-8', errors='replace')
|
||
sys.stderr.reconfigure(encoding='utf-8', errors='replace')
|
||
|
||
# 确定阶段名
|
||
if args.baseline:
|
||
phase_name = "baseline"
|
||
elif args.phase:
|
||
phase_name = args.phase
|
||
else:
|
||
phase_name = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||
|
||
# API 地址覆盖
|
||
api_url = args.api_url or RAG_API_URL
|
||
|
||
# 路径处理
|
||
dataset_path = PROJECT_ROOT / args.dataset
|
||
if not dataset_path.exists():
|
||
logger.error(f"dataset not found: {dataset_path}")
|
||
sys.exit(1)
|
||
|
||
results_dir = PROJECT_ROOT / RESULTS_DIR
|
||
results_dir.mkdir(parents=True, exist_ok=True)
|
||
|
||
# 设置文件日志
|
||
log_path = results_dir / f"eval_{phase_name}.log"
|
||
setup_file_logging(str(log_path))
|
||
|
||
# 输出路径
|
||
if args.output:
|
||
output_path = PROJECT_ROOT / args.output
|
||
else:
|
||
output_path = results_dir / f"phase_{phase_name}.json"
|
||
|
||
report_path = output_path.with_suffix(".md")
|
||
|
||
# 执行评测
|
||
logger.info(f"=== RAG E2E Eval [{phase_name}] ===")
|
||
logger.info(f"dataset: {dataset_path}")
|
||
logger.info(f"API: {api_url}")
|
||
logger.info(f"LLM scoring: {'on' if not args.no_llm else 'off'}")
|
||
|
||
eval_result = evaluate_dataset(
|
||
str(dataset_path),
|
||
use_llm=not args.no_llm,
|
||
api_url=api_url
|
||
)
|
||
|
||
if not eval_result:
|
||
logger.error("eval failed, no results")
|
||
sys.exit(1)
|
||
|
||
# 保存 JSON
|
||
with open(output_path, 'w', encoding='utf-8') as f:
|
||
json.dump(eval_result, f, ensure_ascii=False, indent=2)
|
||
logger.info(f"results saved: {output_path}")
|
||
|
||
# 基线对比
|
||
comparison = None
|
||
baseline_path = results_dir / "phase_baseline.json"
|
||
if phase_name != "baseline" and baseline_path.exists():
|
||
logger.info("comparing with baseline...")
|
||
comparison = compare_results(str(baseline_path), str(output_path))
|
||
regressions = comparison.get("regressions", [])
|
||
if regressions:
|
||
logger.warning(f"found {len(regressions)} regression(s)!")
|
||
else:
|
||
logger.info("no regressions, optimization is safe.")
|
||
|
||
# 生成报告
|
||
report = generate_report(eval_result, comparison, phase_name)
|
||
with open(report_path, 'w', encoding='utf-8') as f:
|
||
f.write(report)
|
||
logger.info(f"report saved: {report_path}")
|
||
|
||
# 控制台摘要(纯 ASCII,避免 Windows 编码问题)
|
||
summary = eval_result.get("summary", {})
|
||
print(f"\n{'='*60}")
|
||
print(f" RAG E2E Eval [{phase_name}] DONE")
|
||
print(f"{'='*60}")
|
||
print(f" Avg Keyword Coverage: {summary.get('avg_keyword_coverage', 0):.2%}")
|
||
print(f" Avg LLM Score: {summary.get('avg_llm_overall', 0):.2f}")
|
||
print(f" Questions: {summary.get('total_questions', 0)}")
|
||
|
||
if comparison:
|
||
sd = comparison["summary_delta"]
|
||
print(f" --- vs Baseline ---")
|
||
print(f" Keyword delta: {sd['keyword_delta']:+.2%}")
|
||
print(f" LLM delta: {sd['llm_delta']:+.2f}")
|
||
regressions = comparison.get("regressions", [])
|
||
if regressions:
|
||
print(f" !! Regressions: {len(regressions)}")
|
||
for r in regressions:
|
||
print(f" - {r['id']}: kw {r['delta_kw']:+.0%}")
|
||
|
||
print(f" JSON: {output_path}")
|
||
print(f" Report: {report_path}")
|
||
print(f" Log: {log_path}")
|
||
print(f"{'='*60}\n")
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|