#!/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()