init: RAG 知识库服务初始提交
- 后端 API(Flask + Gunicorn) - RAG 引擎(混合检索 + 云端 Reranker + 引用溯源) - 文档解析(MinerU + 多格式支持) - Docker 生产部署配置 - 排除前端项目、敏感配置、模型文件
This commit is contained in:
670
scripts/eval_e2e.py
Normal file
670
scripts/eval_e2e.py
Normal file
@@ -0,0 +1,670 @@
|
||||
#!/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 = "data/eval/eval_dataset.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"
|
||||
}
|
||||
# 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 格式返回,不要包含其他内容:
|
||||
{{"accuracy": <分数>, "completeness": <分数>, "relevance": <分数>, "fluency": <分数>, "overall": <总分>}}"""
|
||||
|
||||
try:
|
||||
response = client.chat.completions.create(
|
||||
model=DASHSCOPE_MODEL,
|
||||
messages=[{"role": "user", "content": prompt}],
|
||||
temperature=0.1,
|
||||
max_tokens=300
|
||||
)
|
||||
text = response.choices[0].message.content.strip()
|
||||
|
||||
# 提取 JSON
|
||||
json_match = re.search(r'\{[^}]+\}', text)
|
||||
if json_match:
|
||||
scores = json.loads(json_match.group())
|
||||
# 归一化到 0-1
|
||||
for k in scores:
|
||||
scores[k] = round(min(10, max(0, scores[k])) / 10.0, 4)
|
||||
return scores
|
||||
else:
|
||||
logger.warning(f"LLM 返回内容无法解析 JSON: {text[:100]}")
|
||||
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("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()
|
||||
Reference in New Issue
Block a user