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>
This commit is contained in:
@@ -0,0 +1,348 @@
|
||||
#!/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()
|
||||
Reference in New Issue
Block a user