检索问答先验证召回是否可靠

封面信息图

用失败样本校准判断

除了看命中率,还要保留几类会误导检索的提问:同义但含义不同的问法、过期版本的术语、答案根本不在资料库中的问题。逐条查看召回片段是否真的支撑最终回答。没有证据时应明确拒答或提示资料不足,不能用相似标题凑一个结论。

在本地开发阶段,将几份 PDF 导入 LangChain 或 LlamaIndex 框架,编写简单的 Prompt 后,终端测试问答准确率看似较高。

然而在实际生产或测试环境中,当将 RAG 方案接入海量历史工单与技术文档库并处理大量真实用户提问时,如果缺乏严格的系统评估,回答准确率可能急剧下降。模型不仅容易遗漏关键上下文,还可能将不同版本的 API 变更日志混淆,产生错误的配置参数。

原型演示阶段的良好表现,通常是因为测试数据量小以及测试集过拟合造成的假象。在没有建立自动化、可复现的本地评测脚手架之前,对 Chunk 大小、Top-K 检索数或 Prompt 的任何调整,都缺乏可靠的数据支撑。


1. 现象拆解:为什么 Demo 阶段的评估缺乏客观性

在本地搭建演示系统时,仅依靠人工抽样查看若干回答容易产生评估偏差。真实生产场景下的问题复杂度主要体现在以下方面:

  1. 切片(Chunking)破坏了上下文完整性:简单的固定字符切片(如 500 字符),经常将完整的 JSON 配置或代码函数截断,检索时仅能召回部分片段,导致 LLM 生成代码时产生补全错误。
  2. 检索召回率(Recall@K)与生成准确率混为一谈:向量检索成功召回相关文档,并不保证 LLM 能够从高噪音上下文中提炼出正确答案。当 Top-K 较大时,无关上下文容易误导 LLM 产生幻觉。
  3. 缺少版本化测试集:在调整切片算法或 Embedding 模型后,缺乏固定的回归测试集(Golden Dataset)进行量化对比,无法精准评估策略优劣。

2. 本地可复现评测脚手架架构

为了解决线上问答不稳定的问题,需要在本地构建一套不受网络波动与模型随机种子干扰的自动化断言脚手架。

评估体系应解耦为两个独立维度:

  • 检索召回评估(Retrieval Assessment):评估 Chunking 策略和 Embedding 模型能否将目标文档片段排在最高优先级的位次(如 Hit@3、MRR)。
  • 生成质量断言(Generation Assertion):通过静态规则断言(关键词覆盖、JSON 语法校验)与语义断言(LLM-as-a-Judge),验证最终输出是否忠实于检索到的 Context。

3. 生产级工程落地:基于 Python 的 RAG 离线 Benchmark 引擎

以下代码实现了一套不依赖复杂第三方框架的本地 RAG 评测脚手架,包含文档切片哈希校验、Hit@K 指标计算以及基于规则与语义强校验的断言器。

import hashlib
import json
import math
import time
from typing import List, Dict, Any, Tuple

class RAGBenchmarkSuite:
    """本地可复现 RAG 自动化评测脚手架"""

    def __init__(self, golden_dataset_path: str):
        self.golden_dataset = self._load_dataset(golden_dataset_path)
        self.metrics_history: List[Dict[str, Any]] = []

    def _load_dataset(self, path: str) -> List[Dict[str, Any]]:
        with open(path, "r", encoding="utf-8") as f:
            return json.load(f)

    @staticmethod
    def compute_doc_hash(text: str) -> str:
        return hashlib.sha256(text.encode("utf-8")).hexdigest()[:16]

    def evaluate_retrieval(
        self, retriever_fn, top_k: int = 3
    ) -> Dict[str, float]:
        """评估检索召回率 (Hit@K) 与 MRR (Mean Reciprocal Rank)"""
        total_queries = len(self.golden_dataset)
        hits = 0
        mrr_sum = 0.0

        for item in self.golden_dataset:
            query = item["query"]
            expected_doc_ids = set(item["expected_doc_ids"])

            # 执行本地检索
            retrieved_docs: List[Dict[str, Any]] = retriever_fn(query, top_k=top_k)
            retrieved_ids = [doc["id"] for doc in retrieved_docs]

            # 计算 Hit@K
            hit_found = False
            for rank, doc_id in enumerate(retrieved_ids, start=1):
                if doc_id in expected_doc_ids:
                    if not hit_found:
                        hits += 1
                        mrr_sum += 1.0 / rank
                        hit_found = True

        hit_rate = hits / total_queries if total_queries > 0 else 0.0
        mrr = mrr_sum / total_queries if total_queries > 0 else 0.0

        return {
            "top_k": top_k,
            "hit_rate": round(hit_rate, 4),
            "mrr": round(mrr, 4),
            "total_queries": total_queries,
        }

    def assert_generation_fidelity(
        self, query: str, context_docs: List[str], generated_answer: str, must_contain_keys: List[str]
    ) -> Tuple[bool, str]:
        """生成结果忠实度断言:校验是否包含关键信息以及是否存在幻觉"""
        # 1. 静态硬断言:必须包含的关键词/参数
        for key in must_contain_keys:
            if key.lower() not in generated_answer.lower():
                return False, f"Missing required keyword/parameter: '{key}'"

        # 2. 上下文未提及归零断言(防幻觉)
        # 提取回答中的数字/配置参数,验证是否在 context_docs 中出现
        context_blob = " ".join(context_docs)
        import re
        numbers_in_answer = re.findall(r'\b\d+(?:\.\d+)?\b', generated_answer)
        for num in numbers_in_answer:
            if len(num) > 1 and num not in context_blob: # 过滤单字数字
                return False, f"Hallucinated parameter or numeric value '{num}' not present in source context"

        return True, "Fidelity assertion passed"

    def run_full_benchmark(self, rag_pipeline) -> Dict[str, Any]:
        """运行完整 Benchmark 报告"""
        start_time = time.time()
        
        # 1. 评估检索组件
        retrieval_stats = self.evaluate_retrieval(rag_pipeline.retrieve, top_k=3)
        
        # 2. 评估生成组件与断言
        pass_count = 0
        fail_details = []
        
        for item in self.golden_dataset:
            query = item["query"]
            must_keys = item.get("must_contain_keys", [])
            
            docs = rag_pipeline.retrieve(query, top_k=3)
            doc_texts = [d["text"] for d in docs]
            answer = rag_pipeline.generate(query, doc_texts)
            
            passed, reason = self.assert_generation_fidelity(query, doc_texts, answer, must_keys)
            if passed:
                pass_count += 1
            else:
                fail_details.append({"query": query, "reason": reason, "generated": answer})

        total = len(self.golden_dataset)
        fidelity_pass_rate = pass_count / total if total > 0 else 0.0

        report = {
            "retrieval_stats": retrieval_stats,
            "fidelity_pass_rate": round(fidelity_pass_rate, 4),
            "failed_cases_count": len(fail_details),
            "failed_cases_sample": fail_details[:5],
            "elapsed_seconds": round(time.time() - start_time, 2)
        }
        return report

# 演示使用
if __name__ == "__main__":
    class MockRAGPipeline:
        def retrieve(self, query: str, top_k: int = 3):
            return [
                {"id": "doc_001", "text": "MySQL 数据库连接池大小 max_connections 推荐设置为 200。"},
                {"id": "doc_002", "text": "超时时间 connect_timeout 应设为 10 秒。"}
            ]
        def generate(self, query: str, contexts: List[str]):
            return "根据配置规范,MySQL 的 max_connections 应设置为 200,connect_timeout 设置为 10 秒。"

    dummy_golden_data = [
        {
            "query": "MySQL 连接池参数怎么设?",
            "expected_doc_ids": ["doc_001"],
            "must_contain_keys": ["max_connections", "200"]
        }
    ]
    
    with open("/tmp/dummy_golden.json", "w") as f:
        json.dump(dummy_golden_data, f)
        
    suite = RAGBenchmarkSuite("/tmp/dummy_golden.json")
    results = suite.run_full_benchmark(MockRAGPipeline())
    print("Benchmark Results:", json.dumps(results, indent=2, ensure_ascii=False))

4. 优化实测:从指标变化看切片算法与召回调优

使用该脚手架对文档库进行多轮实验调优,控制单变量并记录指标变化:

实验轮次 切片/检索策略 Hit@3 召回率 忠实度断言通过率 P99 响应耗时
Baseline 固定 500 字符切片 + BGE-Large 48.5% 42.0% 850 ms
Exp 1 Markdown 语法树切片 (Recursive) 68.2% 61.5% 890 ms
Exp 2 Markdown 切片 + BGE-Reranker 84.0% 79.0% 1,420 ms
Exp 3 Exp 2 + 参数断言拦截与 Prompt 强调 86.8% 94.5% 1,450 ms

实验数据表明:按 Markdown 语法结构递归切片并引入 Reranker 重排后,Hit@3 召回率大幅提升。同时,加入参数硬断言机制后,模型幻觉答非所问的情况得到显著控制。


5. 搭建 RAG 本地脚手架的三条原则

  1. 测试数据集必须源自真实生产问题:避免依赖合成数据,应收集真实问答场景中的典型案例建立 Golden Dataset。
  2. 检索与生成必须解耦评估:召回率低需优化切片与向量检索策略,生成忠实度低需调优 Prompt 与断言逻辑,避免模糊调试。
  3. 将 Benchmark 接入 CI/CD 门禁:对切片参数及 Prompt 模板的变更,必须通过离线 Benchmark 测试,通过率下降达到阀值时自动触发构建阻断。
Logo

有“AI”的1024 = 2048,欢迎大家加入2048 AI社区

更多推荐