Files
law-kb/scripts/eval_retrieval.py
T
freedak 641e33b834 feat: QYLAW 法律法规知识库
- 语义检索(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>
2026-08-07 14:55:25 +08:00

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()