641e33b834
- 语义检索(FAISS + embedding)+ 精确查找(法规名+条号) - RAG 问答(SSE 流式,支持 thinking 折叠显示) - 法规浏览(原文阅读) - 历史记录(检索+对话持久化到 SQLite) - 设置页(系统提示词/模板/LLM 参数可配置) - 检索质量评估脚本 Generated with [Devin](https://devin.ai) Co-Authored-By: Devin <158243242+devin-ai-integration[bot]@users.noreply.github.com>
349 lines
12 KiB
Python
349 lines
12 KiB
Python
#!/usr/bin/env python3
|
|
"""检索质量评估脚本
|
|
|
|
从 metadata.db 随机抽取条款,自动生成测试 query,
|
|
调用 /api/search 接口,计算 Recall@K、MRR、NDCG 等指标。
|
|
|
|
用法:
|
|
python3 scripts/eval_retrieval.py [--sample 500] [--top-k 10] [--api http://localhost:8090]
|
|
|
|
两种测试模式:
|
|
1. 结构化 query:"法规名 条号"(如"郑州市劳动用工条例 第三十二条")
|
|
→ 测试精确查找模式(keyword)和语义检索模式(semantic)的命中差异
|
|
2. 内容 query:条款内容前 50 字
|
|
→ 测试语义检索是否能通过条文片段找到原文
|
|
"""
|
|
import argparse
|
|
import asyncio
|
|
import json
|
|
import math
|
|
import random
|
|
import sqlite3
|
|
import sys
|
|
import time
|
|
from pathlib import Path
|
|
from typing import List, Dict, Tuple
|
|
|
|
# 添加项目根目录到 path
|
|
sys.path.insert(0, str(Path(__file__).parent.parent))
|
|
|
|
import httpx
|
|
|
|
|
|
def build_eval_dataset(db_path: str, sample_size: int = 500) -> List[Dict]:
|
|
"""从数据库随机抽取条款,构建评估数据集
|
|
|
|
Returns:
|
|
[{"clause_id", "law_name", "clause_no", "content", "law_id"}, ...]
|
|
"""
|
|
conn = sqlite3.connect(db_path)
|
|
conn.row_factory = sqlite3.Row
|
|
|
|
# 随机抽取条款(排除内容过短的)
|
|
rows = conn.execute(
|
|
"""
|
|
SELECT c.id, c.law_id, c.clause_no, c.content,
|
|
l.name as law_name, l.category, l.province
|
|
FROM clauses c JOIN laws l ON c.law_id = l.id
|
|
WHERE length(c.content) > 20 AND c.clause_no IS NOT NULL
|
|
ORDER BY RANDOM() LIMIT ?
|
|
""",
|
|
(sample_size,),
|
|
).fetchall()
|
|
conn.close()
|
|
|
|
dataset = []
|
|
for r in rows:
|
|
dataset.append({
|
|
"clause_id": r["id"],
|
|
"law_id": r["law_id"],
|
|
"law_name": r["law_name"],
|
|
"clause_no": r["clause_no"],
|
|
"content": r["content"],
|
|
"category": r["category"],
|
|
"province": r["province"],
|
|
})
|
|
return dataset
|
|
|
|
|
|
def generate_queries(item: Dict) -> Dict[str, str]:
|
|
"""为一个条款生成多种 query
|
|
|
|
Returns:
|
|
{"structured": "法规名 条号", "content_snippet": "条款前50字"}
|
|
"""
|
|
# 结构化 query:法规名 + 条号
|
|
structured = f"{item['law_name']} {item['clause_no']}"
|
|
|
|
# 内容 query:条款内容前 50 字(去掉换行)
|
|
content_clean = item["content"].replace("\n", " ").strip()
|
|
snippet = content_clean[:50]
|
|
|
|
return {
|
|
"structured": structured,
|
|
"content_snippet": snippet,
|
|
}
|
|
|
|
|
|
async def search_api(
|
|
client: httpx.AsyncClient,
|
|
api_base: str,
|
|
query: str,
|
|
mode: str = "semantic",
|
|
top_k: int = 10,
|
|
) -> List[Dict]:
|
|
"""调用检索 API
|
|
|
|
Returns:
|
|
结果列表,每项含 clause_id/law_name/clause_no/score 等
|
|
"""
|
|
params = {
|
|
"query": query,
|
|
"mode": mode,
|
|
"top_k": top_k,
|
|
"page": 1,
|
|
"page_size": top_k,
|
|
}
|
|
try:
|
|
resp = await client.get(f"{api_base}/api/search", params=params, timeout=30)
|
|
resp.raise_for_status()
|
|
data = resp.json()
|
|
results = data.get("data", {}).get("results", [])
|
|
# 补充 clause_id(检索 API 返回的字段名)
|
|
for r in results:
|
|
if "clause_id" not in r:
|
|
r["clause_id"] = r.get("id")
|
|
return results
|
|
except Exception as e:
|
|
print(f" [ERROR] 检索失败: {e}")
|
|
return []
|
|
|
|
|
|
def compute_recall_at_k(results: List[Dict], gt_clause_id: int, k: int) -> float:
|
|
"""Recall@K:ground truth 是否在 Top-K 结果中"""
|
|
top_k = results[:k]
|
|
for r in top_k:
|
|
if r.get("clause_id") == gt_clause_id:
|
|
return 1.0
|
|
return 0.0
|
|
|
|
|
|
def compute_mrr(results: List[Dict], gt_clause_id: int) -> float:
|
|
"""MRR:第一个相关文档的排名倒数"""
|
|
for i, r in enumerate(results, 1):
|
|
if r.get("clause_id") == gt_clause_id:
|
|
return 1.0 / i
|
|
return 0.0
|
|
|
|
|
|
def compute_ndcg_at_k(results: List[Dict], gt_clause_id: int, k: int) -> float:
|
|
"""NDCG@K:二值相关性(命中=1,未命中=0)"""
|
|
dcg = 0.0
|
|
for i, r in enumerate(results[:k], 1):
|
|
if r.get("clause_id") == gt_clause_id:
|
|
dcg = 1.0 / math.log2(i + 1)
|
|
break
|
|
# IDCG:理想情况下相关文档排第 1
|
|
idcg = 1.0 / math.log2(2) # = 1.0
|
|
return dcg / idcg if idcg > 0 else 0.0
|
|
|
|
|
|
def compute_hit_rate(results: List[Dict], gt_clause_id: int) -> float:
|
|
"""Hit Rate:是否至少命中一个相关文档"""
|
|
for r in results:
|
|
if r.get("clause_id") == gt_clause_id:
|
|
return 1.0
|
|
return 0.0
|
|
|
|
|
|
async def run_eval(
|
|
dataset: List[Dict],
|
|
api_base: str,
|
|
top_k: int = 10,
|
|
modes: List[str] = None,
|
|
query_types: List[str] = None,
|
|
) -> Dict:
|
|
"""运行评估
|
|
|
|
Args:
|
|
dataset: 评估数据集
|
|
api_base: API 地址
|
|
top_k: Top-K
|
|
modes: 检索模式列表 ["semantic", "keyword"]
|
|
query_types: query 类型列表 ["structured", "content_snippet"]
|
|
|
|
Returns:
|
|
评估结果报告
|
|
"""
|
|
if modes is None:
|
|
modes = ["semantic", "keyword"]
|
|
if query_types is None:
|
|
query_types = ["structured", "content_snippet"]
|
|
|
|
report = {
|
|
"sample_size": len(dataset),
|
|
"top_k": top_k,
|
|
"modes": modes,
|
|
"query_types": query_types,
|
|
"timestamp": time.strftime("%Y-%m-%d %H:%M:%S"),
|
|
"results": {},
|
|
}
|
|
|
|
async with httpx.AsyncClient() as client:
|
|
for mode in modes:
|
|
for qtype in query_types:
|
|
key = f"{mode}_{qtype}"
|
|
print(f"\n=== 评估: mode={mode}, query_type={qtype} ===")
|
|
|
|
recalls = []
|
|
mrrs = []
|
|
ndcgs = []
|
|
hit_rates = []
|
|
latencies = []
|
|
errors = 0
|
|
|
|
for i, item in enumerate(dataset):
|
|
queries = generate_queries(item)
|
|
query = queries.get(qtype, "")
|
|
if not query:
|
|
continue
|
|
|
|
gt_id = item["clause_id"]
|
|
|
|
t0 = time.time()
|
|
results = await search_api(client, api_base, query, mode, top_k)
|
|
latency = (time.time() - t0) * 1000
|
|
latencies.append(latency)
|
|
|
|
if not results:
|
|
errors += 1
|
|
recalls.append(0.0)
|
|
mrrs.append(0.0)
|
|
ndcgs.append(0.0)
|
|
hit_rates.append(0.0)
|
|
else:
|
|
recalls.append(compute_recall_at_k(results, gt_id, top_k))
|
|
mrrs.append(compute_mrr(results, gt_id))
|
|
ndcgs.append(compute_ndcg_at_k(results, gt_id, top_k))
|
|
hit_rates.append(compute_hit_rate(results, gt_id))
|
|
|
|
# 进度
|
|
if (i + 1) % 50 == 0:
|
|
avg_recall = sum(recalls) / len(recalls)
|
|
print(f" 进度: {i+1}/{len(dataset)} | Recall@{top_k}={avg_recall:.3f}")
|
|
|
|
n = len(recalls)
|
|
report["results"][key] = {
|
|
"mode": mode,
|
|
"query_type": qtype,
|
|
"count": n,
|
|
"errors": errors,
|
|
"recall_at_k": sum(recalls) / n if n else 0,
|
|
"mrr": sum(mrrs) / n if n else 0,
|
|
"ndcg_at_k": sum(ndcgs) / n if n else 0,
|
|
"hit_rate": sum(hit_rates) / n if n else 0,
|
|
"avg_latency_ms": sum(latencies) / len(latencies) if latencies else 0,
|
|
"p95_latency_ms": sorted(latencies)[int(len(latencies) * 0.95)] if latencies else 0,
|
|
}
|
|
|
|
r = report["results"][key]
|
|
print(f" 结果: Recall@{top_k}={r['recall_at_k']:.3f} | MRR={r['mrr']:.3f} | "
|
|
f"NDCG@{top_k}={r['ndcg_at_k']:.3f} | HitRate={r['hit_rate']:.3f} | "
|
|
f"avg={r['avg_latency_ms']:.0f}ms | errors={errors}")
|
|
|
|
return report
|
|
|
|
|
|
def print_report(report: Dict):
|
|
"""打印评估报告"""
|
|
print("\n" + "=" * 80)
|
|
print("检索质量评估报告")
|
|
print("=" * 80)
|
|
print(f"时间: {report['timestamp']}")
|
|
print(f"样本数: {report['sample_size']}")
|
|
print(f"Top-K: {report['top_k']}")
|
|
print()
|
|
|
|
# 表格输出
|
|
print(f"{'模式':<12} {'Query类型':<18} {'Recall@K':>10} {'MRR':>10} {'NDCG@K':>10} {'HitRate':>10} {'avg(ms)':>10} {'错误':>6}")
|
|
print("-" * 96)
|
|
for key, r in report["results"].items():
|
|
print(f"{r['mode']:<12} {r['query_type']:<18} "
|
|
f"{r['recall_at_k']:>10.3f} {r['mrr']:>10.3f} {r['ndcg_at_k']:>10.3f} "
|
|
f"{r['hit_rate']:>10.3f} {r['avg_latency_ms']:>10.0f} {r['errors']:>6}")
|
|
|
|
print()
|
|
# 分析建议
|
|
sem_struct = report["results"].get("semantic_structured", {})
|
|
kw_struct = report["results"].get("keyword_structured", {})
|
|
sem_content = report["results"].get("semantic_content_snippet", {})
|
|
|
|
print("分析:")
|
|
if sem_struct and kw_struct:
|
|
if kw_struct["recall_at_k"] > sem_struct["recall_at_k"] + 0.1:
|
|
print(f" - 结构化查询(法规名+条号):精确查找(Recall={kw_struct['recall_at_k']:.3f})"
|
|
f" 显著优于语义检索(Recall={sem_struct['recall_at_k']:.3f})")
|
|
print(f" → 建议:用户输入法规名+条号时,自动切换到精确查找模式")
|
|
else:
|
|
print(f" - 结构化查询:语义检索(Recall={sem_struct['recall_at_k']:.3f})"
|
|
f" 与精确查找(Recall={kw_struct['recall_at_k']:.3f})相当")
|
|
|
|
if sem_content:
|
|
r = sem_content
|
|
if r["recall_at_k"] > 0.8:
|
|
print(f" - 内容片段查询:语义检索表现优秀(Recall={r['recall_at_k']:.3f})")
|
|
elif r["recall_at_k"] > 0.5:
|
|
print(f" - 内容片段查询:语义检索表现一般(Recall={r['recall_at_k']:.3f}),有提升空间")
|
|
else:
|
|
print(f" - 内容片段查询:语义检索表现较差(Recall={r['recall_at_k']:.3f}),需检查 embedding 模型")
|
|
|
|
print()
|
|
|
|
|
|
def main():
|
|
parser = argparse.ArgumentParser(description="检索质量评估")
|
|
parser.add_argument("--sample", type=int, default=500, help="抽样数量(默认 500)")
|
|
parser.add_argument("--top-k", type=int, default=10, help="Top-K(默认 10)")
|
|
parser.add_argument("--api", type=str, default="http://localhost:8090", help="API 地址")
|
|
parser.add_argument("--db", type=str, default="/data/law-kb-data/metadata.db", help="SQLite 路径")
|
|
parser.add_argument("--seed", type=int, default=42, help="随机种子(默认 42)")
|
|
parser.add_argument("--output", type=str, default=None, help="结果输出 JSON 文件路径")
|
|
parser.add_argument("--modes", type=str, nargs="+", default=["semantic", "keyword"],
|
|
help="检索模式(默认 semantic keyword)")
|
|
parser.add_argument("--query-types", type=str, nargs="+",
|
|
default=["structured", "content_snippet"],
|
|
help="query 类型(默认 structured content_snippet)")
|
|
args = parser.parse_args()
|
|
|
|
random.seed(args.seed)
|
|
|
|
print(f"构建评估数据集(抽样 {args.sample} 条)...")
|
|
dataset = build_eval_dataset(args.db, args.sample)
|
|
print(f"数据集大小: {len(dataset)}")
|
|
|
|
if not dataset:
|
|
print("错误:数据集为空,检查数据库路径")
|
|
sys.exit(1)
|
|
|
|
# 运行评估
|
|
report = asyncio.run(run_eval(
|
|
dataset=dataset,
|
|
api_base=args.api,
|
|
top_k=args.top_k,
|
|
modes=args.modes,
|
|
query_types=args.query_types,
|
|
))
|
|
|
|
# 打印报告
|
|
print_report(report)
|
|
|
|
# 保存结果
|
|
if args.output:
|
|
with open(args.output, "w", encoding="utf-8") as f:
|
|
json.dump(report, f, ensure_ascii=False, indent=2)
|
|
print(f"结果已保存到: {args.output}")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|