51dc8dc4f6
- FastAPI + Qdrant + Redis + Ollama 技术栈 - L1→L2→L3→chunk 四层分层检索(dense + sparse RRF 融合) - 文档三级总结与 2.5 级回退 - query 解析路由与分类 - /admin 管理页面
354 lines
15 KiB
Python
354 lines
15 KiB
Python
#!/usr/bin/env python3
|
||
"""分层摘要 RAG 离线评测 harness
|
||
|
||
流程:回归集 → 独立 _eval 后缀集合入库(Ingester)→ 摘要质量指标(Entity Recall /
|
||
幻觉率 / 类目一致性)→ 检索效用指标(Routing F1 / Pruning Loss / Precision@5 /
|
||
Recall@10)→ 平铺 chunks baseline 对比 → Markdown 报告(stdout + report.md)。
|
||
|
||
用法:
|
||
uv run python scripts/eval/run_eval.py [--regression PATH] [--keep-data] [--no-judge]
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import argparse
|
||
import asyncio
|
||
import json
|
||
import sys
|
||
from pathlib import Path
|
||
|
||
# 脚本直运行时把项目根加入 sys.path,保证可以 import app 与 scripts.eval
|
||
_PROJECT_ROOT = Path(__file__).resolve().parents[2]
|
||
if str(_PROJECT_ROOT) not in sys.path:
|
||
sys.path.insert(0, str(_PROJECT_ROOT))
|
||
|
||
import structlog # noqa: E402
|
||
|
||
from app.config import settings # noqa: E402
|
||
from app.core import ingestion as ingestion_mod # noqa: E402
|
||
from app.core import retriever as retriever_mod # noqa: E402
|
||
from app.core.ingestion import Ingester # noqa: E402
|
||
from app.core.retriever import Retriever # noqa: E402
|
||
from app.models.document import DocumentInput # noqa: E402
|
||
from app.models.search import SearchRequest # noqa: E402
|
||
from app.services import qdrant as qdrant_mod # noqa: E402
|
||
from app.services.ollama import OllamaClient # noqa: E402
|
||
from app.services.qdrant import QdrantService # noqa: E402
|
||
from scripts.eval.judge import hallucination_rate, taxonomy_consistency # noqa: E402
|
||
from scripts.eval.metrics import ( # noqa: E402
|
||
aggregate,
|
||
entity_recall,
|
||
precision_at_k,
|
||
pruning_loss,
|
||
recall_at_k,
|
||
routing_f1,
|
||
)
|
||
|
||
logger = structlog.get_logger()
|
||
|
||
# 评测集合后缀与默认路径
|
||
EVAL_SUFFIX = "_eval"
|
||
_SCRIPT_DIR = Path(__file__).resolve().parent
|
||
DEFAULT_REGRESSION = _SCRIPT_DIR / "regression_set.json"
|
||
REPORT_PATH = _SCRIPT_DIR / "report.md"
|
||
|
||
# 门槛(Spec 规定):低于/高于门槛在报告中标红
|
||
THRESHOLD_L1_ER = 0.85 # L1 Entity Recall ≥ 0.85
|
||
THRESHOLD_L3_ER = 0.9 # L3 Entity Recall ≥ 0.9
|
||
THRESHOLD_HALLUCINATION = 0.02 # 幻觉率 < 2%
|
||
THRESHOLD_PRUNING_LOSS = 0.08 # Pruning Loss < 8%
|
||
|
||
|
||
def _switch_to_eval_collections() -> list[str]:
|
||
"""把 app 内模块级集合名常量整体切换为 _eval 后缀的评测集合
|
||
|
||
QdrantService / Retriever / Ingester 均在各自模块命名空间引用了集合名常量,
|
||
评测脚本统一改写这些模块属性实现集合隔离,不改动 app 源码。
|
||
返回评测集合名列表(用于评测结束后清理)。
|
||
"""
|
||
eval_names = {
|
||
"COLLECTION_L1": f"{qdrant_mod.COLLECTION_L1}{EVAL_SUFFIX}",
|
||
"COLLECTION_L2": f"{qdrant_mod.COLLECTION_L2}{EVAL_SUFFIX}",
|
||
"COLLECTION_L3": f"{qdrant_mod.COLLECTION_L3}{EVAL_SUFFIX}",
|
||
"COLLECTION_CHUNKS": f"{qdrant_mod.COLLECTION_CHUNKS}{EVAL_SUFFIX}",
|
||
}
|
||
for module in (qdrant_mod, retriever_mod, ingestion_mod):
|
||
for name, value in eval_names.items():
|
||
if hasattr(module, name):
|
||
setattr(module, name, value)
|
||
qdrant_mod.ALL_COLLECTIONS = tuple(eval_names.values())
|
||
sparse_eval = (eval_names["COLLECTION_L1"], eval_names["COLLECTION_CHUNKS"])
|
||
qdrant_mod.SPARSE_COLLECTIONS = sparse_eval
|
||
retriever_mod.SPARSE_COLLECTIONS = sparse_eval
|
||
# upsert_nodes 按集合名校验层级前缀,需同步替换
|
||
qdrant_mod._NODE_ID_PREFIX = {eval_names["COLLECTION_L2"]: "l2", eval_names["COLLECTION_L3"]: "l3"}
|
||
return list(eval_names.values())
|
||
|
||
|
||
async def _check_services(qdrant: QdrantService, ollama: OllamaClient) -> str | None:
|
||
"""检查 Qdrant / Ollama 连通性,返回错误消息(None 表示正常)"""
|
||
try:
|
||
await qdrant.client.get_collections()
|
||
except Exception as exc:
|
||
return f"无法连接 Qdrant({settings.qdrant_host}:{settings.qdrant_port}):{exc}"
|
||
if not await ollama.is_available():
|
||
return f"无法连接 Ollama({settings.ollama_base_url}),请确认服务已启动(入库与评测均依赖 Ollama)"
|
||
return None
|
||
|
||
|
||
async def _l1_candidate_doc_ids(retriever: Retriever, query: str) -> set[str]:
|
||
"""轻量复现 Retriever 的 L1 路由层:embed → L1 集合 top-N → 候选 doc_id 集合
|
||
|
||
用于计算 Pruning Loss 与 Routing F1;不套类目过滤,度量纯 L1 向量召回。
|
||
"""
|
||
dense = (await retriever.embedding.embed([query]))[0]
|
||
sparse = retriever.sparse_encoder.encode(query) if settings.sparse_enabled else None
|
||
# 复用 Retriever 内部检索方法(hybrid/dense 按集合能力自动选择)
|
||
hits = await retriever._search_collection(retriever_mod.COLLECTION_L1, dense, sparse, settings.l1_doc_top_n, None)
|
||
return {(p.payload or {}).get("doc_id", "") for p in hits} - {""}
|
||
|
||
|
||
async def _baseline_doc_ids(retriever: Retriever, query: str) -> list[str]:
|
||
"""平铺 baseline:直接对 chunks 集合做 hybrid top-k(无路由无剪枝),返回命中 doc_id 列表"""
|
||
dense = (await retriever.embedding.embed([query]))[0]
|
||
sparse = retriever.sparse_encoder.encode(query) if settings.sparse_enabled else None
|
||
points = await retriever._search_collection(
|
||
retriever_mod.COLLECTION_CHUNKS, dense, sparse, settings.retrieval_top_k, None
|
||
)
|
||
return [(p.payload or {}).get("doc_id", "") for p in points]
|
||
|
||
|
||
def _mark(ok: bool) -> str:
|
||
"""门槛判定标记"""
|
||
return "✅" if ok else "❌"
|
||
|
||
|
||
def _fmt(value: float | None, percent: bool = False) -> str:
|
||
"""数值格式化;None 表示未评测(judge 跳过)"""
|
||
if value is None:
|
||
return "N/A"
|
||
return f"{value:.1%}" if percent else f"{value:.4f}"
|
||
|
||
|
||
def _build_report(
|
||
regression_path: Path,
|
||
doc_records: list[dict],
|
||
query_summary: dict,
|
||
hier_metrics: dict[str, float],
|
||
baseline_metrics: dict[str, float],
|
||
routing: dict[str, float],
|
||
prune_loss: float,
|
||
judge_enabled: bool,
|
||
kept_data: bool,
|
||
) -> str:
|
||
"""组装 Markdown 评测报告"""
|
||
lines: list[str] = [
|
||
"# 分层摘要 RAG 评测报告",
|
||
"",
|
||
f"- 回归集:`{regression_path}`",
|
||
f"- 文档数:{len(doc_records)};query 数:{query_summary['total']}"
|
||
f"(positive {query_summary['positive']} / negative {query_summary['negative']})",
|
||
f"- LLM judge:{'开启' if judge_enabled else '跳过(--no-judge 或 Ollama 不可用)'}",
|
||
f"- 评测集合:`*{EVAL_SUFFIX}`({'保留' if kept_data else '已清理'})",
|
||
"",
|
||
"## 摘要质量(按文档)",
|
||
"",
|
||
"| 文档 | 标题 | golden 类目 | 实际类目 | L1 Entity Recall | L3 Entity Recall | 幻觉率 | 类目一致 |",
|
||
"| --- | --- | --- | --- | --- | --- | --- | --- |",
|
||
]
|
||
for record in doc_records:
|
||
consistency = record.get("taxonomy_consistency")
|
||
lines.append(
|
||
f"| {record['id']} | {record['title']} | {record['golden_category']} | {record['category']} "
|
||
f"| {_fmt(record['entity_recall_l1'])} | {_fmt(record['entity_recall_l3'])} "
|
||
f"| {_fmt(record.get('hallucination_rate'), percent=True)} "
|
||
f"| {('是' if consistency else '否') if consistency is not None else 'N/A'} |"
|
||
)
|
||
|
||
# 汇总与门槛判定
|
||
er_l1 = aggregate([{"v": r["entity_recall_l1"]} for r in doc_records]).get("v", 0.0)
|
||
er_l3 = aggregate([{"v": r["entity_recall_l3"]} for r in doc_records]).get("v", 0.0)
|
||
hall_records = [{"v": r["hallucination_rate"]} for r in doc_records if r.get("hallucination_rate") is not None]
|
||
hall = aggregate(hall_records).get("v") if hall_records else None
|
||
lines += [
|
||
"",
|
||
"## 指标汇总与门槛判定",
|
||
"",
|
||
"| 指标 | 数值 | 门槛 | 判定 |",
|
||
"| --- | --- | --- | --- |",
|
||
f"| L1 Entity Recall | {_fmt(er_l1)} | ≥ {THRESHOLD_L1_ER} | {_mark(er_l1 >= THRESHOLD_L1_ER)} |",
|
||
f"| L3 Entity Recall | {_fmt(er_l3)} | ≥ {THRESHOLD_L3_ER} | {_mark(er_l3 >= THRESHOLD_L3_ER)} |",
|
||
f"| Hallucination Rate | {_fmt(hall, percent=True)} | < {THRESHOLD_HALLUCINATION:.0%} "
|
||
f"| {_mark(hall < THRESHOLD_HALLUCINATION) if hall is not None else 'N/A'} |",
|
||
f"| Pruning Loss | {_fmt(prune_loss, percent=True)} | < {THRESHOLD_PRUNING_LOSS:.0%} "
|
||
f"| {_mark(prune_loss < THRESHOLD_PRUNING_LOSS)} |",
|
||
"",
|
||
"## 检索效用(hierarchical vs 平铺 baseline)",
|
||
"",
|
||
"| 指标 | 分层检索 | 平铺 baseline |",
|
||
"| --- | --- | --- |",
|
||
f"| Precision@5 | {_fmt(hier_metrics.get('precision@5', 0.0))} "
|
||
f"| {_fmt(baseline_metrics.get('precision@5', 0.0))} |",
|
||
f"| Recall@10 | {_fmt(hier_metrics.get('recall@10', 0.0))} | {_fmt(baseline_metrics.get('recall@10', 0.0))} |",
|
||
"",
|
||
"### 路由层(L1)",
|
||
"",
|
||
f"- Routing Precision:{_fmt(routing['precision'])}",
|
||
f"- Routing Recall:{_fmt(routing['recall'])}",
|
||
f"- Routing F1:{_fmt(routing['f1'])}",
|
||
f"- Pruning Loss:{_fmt(prune_loss, percent=True)}({_mark(prune_loss < THRESHOLD_PRUNING_LOSS)})",
|
||
"",
|
||
"### negative query",
|
||
"",
|
||
f"- 误中(返回了任意结果):{query_summary['negative_false_alarm']} / {query_summary['negative']}",
|
||
"",
|
||
]
|
||
return "\n".join(lines)
|
||
|
||
|
||
async def run(regression_path: Path, keep_data: bool, judge_enabled: bool) -> int:
|
||
"""评测主流程,返回进程退出码"""
|
||
data = json.loads(regression_path.read_text(encoding="utf-8"))
|
||
documents: list[dict] = data["documents"]
|
||
queries: list[dict] = data["queries"]
|
||
|
||
eval_collections = _switch_to_eval_collections()
|
||
qdrant = QdrantService()
|
||
ollama = OllamaClient()
|
||
|
||
# 连通性检查:失败给出中文提示并以退出码 2 结束
|
||
error = await _check_services(qdrant, ollama)
|
||
if error:
|
||
print(f"错误:{error}", file=sys.stderr)
|
||
return 2
|
||
|
||
await qdrant.ensure_collections()
|
||
logger.info("评测集合就绪", collections=eval_collections)
|
||
|
||
ingester = Ingester(qdrant=qdrant)
|
||
retriever = Retriever(qdrant=qdrant)
|
||
|
||
id_map: dict[str, str] = {} # 回归集 doc id -> 实际入库 doc_id
|
||
doc_records: list[dict] = []
|
||
hier_records: list[dict[str, float]] = []
|
||
baseline_records: list[dict[str, float]] = []
|
||
pruned_cases: list[bool] = []
|
||
golden_sets: list[set[str]] = []
|
||
routed_sets: list[set[str]] = []
|
||
negative_total = 0
|
||
negative_false_alarm = 0
|
||
|
||
try:
|
||
# 1. 逐篇入库并计算摘要质量指标
|
||
for doc in documents:
|
||
result = await ingester.ingest(DocumentInput(title=doc["title"], text=doc["text"]))
|
||
id_map[doc["id"]] = result.document_id
|
||
logger.info("文档入库完成", id=doc["id"], doc_id=result.document_id, category=result.category)
|
||
|
||
record: dict = {
|
||
"id": doc["id"],
|
||
"title": doc["title"],
|
||
"golden_category": doc["golden_category"],
|
||
"category": result.category,
|
||
"entity_recall_l1": entity_recall(doc["text"], result.summary.l1_summary),
|
||
"entity_recall_l3": entity_recall(doc["text"], result.summary.l3_content_outline),
|
||
"hallucination_rate": None,
|
||
"taxonomy_consistency": None,
|
||
}
|
||
if judge_enabled:
|
||
record["hallucination_rate"] = await hallucination_rate(result.summary.l1_summary, doc["text"], ollama)
|
||
record["taxonomy_consistency"] = await taxonomy_consistency(
|
||
result.summary.l1_summary, result.category, ollama
|
||
)
|
||
doc_records.append(record)
|
||
|
||
# 2. 逐 query 计算检索效用指标(分层检索 + 轻量 L1 路由层 + 平铺 baseline)
|
||
for query in queries:
|
||
query_text = query["query"]
|
||
golden_ids = {id_map[query["golden_doc_id"]]} if query.get("golden_doc_id") else set()
|
||
|
||
response = await retriever.search(SearchRequest(query=query_text))
|
||
hit_doc_ids = [hit.doc_id for hit in response.hits]
|
||
l1_candidates = await _l1_candidate_doc_ids(retriever, query_text)
|
||
|
||
if query["type"] == "positive" and golden_ids:
|
||
pruned_cases.append(not golden_ids & l1_candidates)
|
||
golden_sets.append(golden_ids)
|
||
routed_sets.append(l1_candidates)
|
||
hier_records.append(
|
||
{
|
||
"precision@5": precision_at_k(hit_doc_ids, golden_ids, 5),
|
||
"recall@10": recall_at_k(hit_doc_ids, golden_ids, 10),
|
||
}
|
||
)
|
||
baseline_ids = await _baseline_doc_ids(retriever, query_text)
|
||
baseline_records.append(
|
||
{
|
||
"precision@5": precision_at_k(baseline_ids, golden_ids, 5),
|
||
"recall@10": recall_at_k(baseline_ids, golden_ids, 10),
|
||
}
|
||
)
|
||
else:
|
||
# negative query:期望无结果,任何返回均计为误中
|
||
negative_total += 1
|
||
if hit_doc_ids:
|
||
negative_false_alarm += 1
|
||
finally:
|
||
# 评测集合清理(--keep-data 时保留)
|
||
if keep_data:
|
||
logger.info("--keep-data 生效,保留评测集合", collections=eval_collections)
|
||
else:
|
||
for collection in eval_collections:
|
||
try:
|
||
await qdrant.client.delete_collection(collection)
|
||
except Exception as exc:
|
||
logger.warning("评测集合清理失败", collection=collection, error=str(exc))
|
||
logger.info("评测集合已清理", collections=eval_collections)
|
||
|
||
# 3. 汇总并输出报告
|
||
prune_loss = pruning_loss(pruned_cases)
|
||
routing = routing_f1(golden_sets, routed_sets)
|
||
hier_metrics = aggregate(hier_records)
|
||
baseline_metrics = aggregate(baseline_records)
|
||
query_summary = {
|
||
"total": len(queries),
|
||
"positive": len(queries) - negative_total,
|
||
"negative": negative_total,
|
||
"negative_false_alarm": negative_false_alarm,
|
||
}
|
||
report = _build_report(
|
||
regression_path,
|
||
doc_records,
|
||
query_summary,
|
||
hier_metrics,
|
||
baseline_metrics,
|
||
routing,
|
||
prune_loss,
|
||
judge_enabled,
|
||
keep_data,
|
||
)
|
||
print(report)
|
||
REPORT_PATH.write_text(report + "\n", encoding="utf-8")
|
||
logger.info("评测报告已写入", path=str(REPORT_PATH))
|
||
return 0
|
||
|
||
|
||
def main() -> int:
|
||
parser = argparse.ArgumentParser(
|
||
description="分层摘要 RAG 离线评测:回归集入库 → 摘要质量/检索效用指标 → 平铺 baseline 对比 → Markdown 报告",
|
||
)
|
||
parser.add_argument(
|
||
"--regression",
|
||
type=Path,
|
||
default=DEFAULT_REGRESSION,
|
||
help=f"回归集 JSON 路径(默认 {DEFAULT_REGRESSION})",
|
||
)
|
||
parser.add_argument("--keep-data", action="store_true", help="评测结束后保留 _eval 集合(默认删除)")
|
||
parser.add_argument("--no-judge", action="store_true", help="跳过 LLM-as-judge 指标(幻觉率 / 类目一致性)")
|
||
args = parser.parse_args()
|
||
return asyncio.run(run(args.regression.resolve(), args.keep_data, not args.no_judge))
|
||
|
||
|
||
if __name__ == "__main__":
|
||
sys.exit(main())
|