|
@@ -0,0 +1,1608 @@
|
|
|
|
|
+#!/usr/bin/env python3
|
|
|
|
|
+"""
|
|
|
|
|
+RAG 全链路优化系统 — 员工手册智能问答
|
|
|
|
|
+============================================
|
|
|
|
|
+8 项核心优化:
|
|
|
|
|
+ 1. 使用 XX科技有限公司员工手册.md 作为测试文档
|
|
|
|
|
+ 2. 检索器封装为 Agent 工具,自动判断是否需要检索
|
|
|
|
|
+ 3. 使用本地 BGE-Large-Zh-v1.5 embedding 模型 (1024d)
|
|
|
|
|
+ 4. 语义文本切块 (SemanticChunker, percentile=60)
|
|
|
|
|
+ 5. 多路召回: BM25 / 向量 / 混合(Ensemble) / 多查询(MultiQuery)
|
|
|
|
|
+ 6. 查询增强: 查询重写 + HyDE
|
|
|
|
|
+ 7. BGE-Reranker 精排
|
|
|
|
|
+ 8. RAG 效果评估: 测试集生成 → Hit@K/MRR/Recall → 幻觉率 → Bad Case 分析
|
|
|
|
|
+
|
|
|
|
|
+运行方式: python copy_rag_chain.py
|
|
|
|
|
+"""
|
|
|
|
|
+
|
|
|
|
|
+# ============================================================================
|
|
|
|
|
+# Section 1: 导入 + 路径常量 + 配置常量
|
|
|
|
|
+# ============================================================================
|
|
|
|
|
+from __future__ import annotations
|
|
|
|
|
+
|
|
|
|
|
+import json
|
|
|
|
|
+import os
|
|
|
|
|
+import re
|
|
|
|
|
+import sys
|
|
|
|
|
+import shutil
|
|
|
|
|
+import warnings
|
|
|
|
|
+from pathlib import Path
|
|
|
|
|
+from typing import Optional, Union
|
|
|
|
|
+from collections import defaultdict
|
|
|
|
|
+
|
|
|
|
|
+# --- 抑制无关日志 ---
|
|
|
|
|
+warnings.filterwarnings("ignore", category=DeprecationWarning)
|
|
|
|
|
+os.environ.setdefault("GRPC_VERBOSITY", "NONE")
|
|
|
|
|
+os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
|
|
|
|
|
+
|
|
|
|
|
+# --- 加载 .env ---
|
|
|
|
|
+from dotenv import load_dotenv
|
|
|
|
|
+load_dotenv(Path(__file__).resolve().parent.parent / ".env")
|
|
|
|
|
+
|
|
|
|
|
+# --- 路径常量 ---
|
|
|
|
|
+_PROJECT_ROOT = Path(__file__).resolve().parent.parent
|
|
|
|
|
+_CURRENT_DIR = Path(__file__).resolve().parent
|
|
|
|
|
+_MODELS_DIR = _PROJECT_ROOT / "models"
|
|
|
|
|
+_BGE_PATH = _MODELS_DIR / "BAAI" / "bge-large-zh-v1.5"
|
|
|
|
|
+_RERANKER_PATH = _MODELS_DIR / "BAAI" / "bge-reranker-base"
|
|
|
|
|
+_DATA_PATH = _PROJECT_ROOT / "XX科技有限公司员工手册.md"
|
|
|
|
|
+_CHROMA_DIR = _CURRENT_DIR / "chroma_db_bge"
|
|
|
|
|
+
|
|
|
|
|
+# --- 配置常量 ---
|
|
|
|
|
+# 各参数取值理由:经过员工手册(~7600 字、62 个条款)实测调优
|
|
|
|
|
+CHUNK_PERCENTILE = 45 # 语义切块阈值:45 在粒度和语义完整间取平衡(值越小切块越细)
|
|
|
|
|
+RETRIEVAL_K = 10 # 召回阶段返回数量:取 10 是精度/速度折中(太小漏召回,太大稀释 Reranker)
|
|
|
|
|
+DEFAULT_K = 4 # 最终给 LLM 的文档数:4 篇足以覆盖一个问题,再多会撑爆上下文窗口
|
|
|
|
|
+ENSEMBLE_WEIGHTS = [0.4, 0.6] # 混合检索权重 [BM25, 向量]:政策文档关键词重要,BM25 提到 0.4
|
|
|
|
|
+MULTI_QUERY_VARIANTS = 3 # 多查询变体数:3 个变体已足够覆盖同义词,再多边际效益递减
|
|
|
|
|
+N_EVAL_QA = 25 # 评估测试集大小:覆盖全部 16 章的最少问题数
|
|
|
|
|
+RRF_K = 60 # RRF 平滑参数:业界标准值,避免排名靠前的文档过度霸占分数
|
|
|
|
|
+
|
|
|
|
|
+# --- 确保根目录在 sys.path 中 ---
|
|
|
|
|
+sys.path.insert(0, str(_PROJECT_ROOT))
|
|
|
|
|
+
|
|
|
|
|
+# ============================================================================
|
|
|
|
|
+# Section 2: 模型初始化
|
|
|
|
|
+# ============================================================================
|
|
|
|
|
+
|
|
|
|
|
+# --- 2a. Embedding 模型 ---
|
|
|
|
|
+
|
|
|
|
|
+class SentenceTransformerEmbeddings:
|
|
|
|
|
+ """将 SentenceTransformer 包装为 LangChain Embeddings 接口。
|
|
|
|
|
+
|
|
|
|
|
+ 支持 BGE 系列模型的 query_prefix 机制:
|
|
|
|
|
+ - embed_query: 查询前加前缀「为这个句子生成表示以用于检索相关文章:」
|
|
|
|
|
+ - embed_documents: 文档不加前缀(BGE 官方推荐用法)
|
|
|
|
|
+ """
|
|
|
|
|
+ def __init__(self, model_path: str, device: str = "cpu",
|
|
|
|
|
+ query_prefix: str = ""):
|
|
|
|
|
+ from sentence_transformers import SentenceTransformer
|
|
|
|
|
+ self._model = SentenceTransformer(model_path, device=device)
|
|
|
|
|
+ self._query_prefix = query_prefix
|
|
|
|
|
+
|
|
|
|
|
+ def embed_documents(self, texts: list[str]) -> list[list[float]]:
|
|
|
|
|
+ # 文档侧不加前缀
|
|
|
|
|
+ embeddings = self._model.encode(
|
|
|
|
|
+ texts,
|
|
|
|
|
+ normalize_embeddings=True,
|
|
|
|
|
+ show_progress_bar=False,
|
|
|
|
|
+ batch_size=32,
|
|
|
|
|
+ )
|
|
|
|
|
+ return embeddings.tolist()
|
|
|
|
|
+
|
|
|
|
|
+ def embed_query(self, text: str) -> list[float]:
|
|
|
|
|
+ # 查询侧加前缀(BGE 模型关键!不加前缀效果掉 10%+)
|
|
|
|
|
+ if self._query_prefix:
|
|
|
|
|
+ text = self._query_prefix + text
|
|
|
|
|
+ embedding = self._model.encode(
|
|
|
|
|
+ [text],
|
|
|
|
|
+ normalize_embeddings=True,
|
|
|
|
|
+ show_progress_bar=False,
|
|
|
|
|
+ )
|
|
|
|
|
+ return embedding[0].tolist()
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+_embedding_model: Optional[SentenceTransformerEmbeddings] = None
|
|
|
|
|
+
|
|
|
|
|
+# BGE 系列模型的 query_prefix(官方推荐,不加的话检索精度明显下降)
|
|
|
|
|
+_BGE_QUERY_PREFIX = "为这个句子生成表示以用于检索相关文章:"
|
|
|
|
|
+
|
|
|
|
|
+def init_embedding_model() -> SentenceTransformerEmbeddings:
|
|
|
|
|
+ """初始化 BGE-Large-Zh-v1.5 embedding 模型(1024 维)。
|
|
|
|
|
+
|
|
|
|
|
+ BGE-Large-Zh-v1.5 是 BAAI 中文优化版,MTEB 中文榜单榜首。
|
|
|
|
|
+ 相比之前的 Corom (768d),维度更高、中文理解更强。
|
|
|
|
|
+ 关键: 查询时必须加 query_prefix,文档不加。
|
|
|
|
|
+ """
|
|
|
|
|
+ global _embedding_model
|
|
|
|
|
+ if _embedding_model is not None:
|
|
|
|
|
+ return _embedding_model
|
|
|
|
|
+ if not _BGE_PATH.exists():
|
|
|
|
|
+ raise FileNotFoundError(f"BGE 模型未找到: {_BGE_PATH}")
|
|
|
|
|
+ print(f"[Embedding] 加载 BGE-Large-Zh-v1.5: {_BGE_PATH}")
|
|
|
|
|
+ _embedding_model = SentenceTransformerEmbeddings(
|
|
|
|
|
+ str(_BGE_PATH), device="cpu",
|
|
|
|
|
+ query_prefix=_BGE_QUERY_PREFIX,
|
|
|
|
|
+ )
|
|
|
|
|
+ print(f"[Embedding] 加载完成,维度: 1024 (含 query_prefix)")
|
|
|
|
|
+ return _embedding_model
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+# --- 2b. LLM 客户端 ---
|
|
|
|
|
+
|
|
|
|
|
+_llm = None
|
|
|
|
|
+
|
|
|
|
|
+def get_llm():
|
|
|
|
|
+ """获取 LLM 客户端(复用根目录 llm_client.py 的 get_llm())。"""
|
|
|
|
|
+ global _llm
|
|
|
|
|
+ if _llm is not None:
|
|
|
|
|
+ return _llm
|
|
|
|
|
+ from llm_client import get_llm as _root_get_llm
|
|
|
|
|
+ _llm = _root_get_llm()
|
|
|
|
|
+ print(f"[LLM] 客户端初始化完成")
|
|
|
|
|
+ return _llm
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+# --- 2c. Reranker 模型 ---
|
|
|
|
|
+
|
|
|
|
|
+_reranker = None
|
|
|
|
|
+
|
|
|
|
|
+def download_reranker() -> Optional[str]:
|
|
|
|
|
+ """从 ModelScope 下载 BGE-Reranker-Base。
|
|
|
|
|
+
|
|
|
|
|
+ 返回模型路径,如果下载失败返回 None。
|
|
|
|
|
+ 已下载则直接返回路径,跳过下载。
|
|
|
|
|
+ """
|
|
|
|
|
+ config_file = _RERANKER_PATH / "config.json"
|
|
|
|
|
+ if config_file.exists():
|
|
|
|
|
+ print(f"[Reranker] 模型已存在: {_RERANKER_PATH}")
|
|
|
|
|
+ return str(_RERANKER_PATH)
|
|
|
|
|
+
|
|
|
|
|
+ print(f"[Reranker] 正在从 ModelScope 下载 BGE-Reranker-Base (~1.1GB)...")
|
|
|
|
|
+ print(f"[Reranker] 目标路径: {_RERANKER_PATH}")
|
|
|
|
|
+ try:
|
|
|
|
|
+ from modelscope import snapshot_download
|
|
|
|
|
+ _RERANKER_PATH.parent.mkdir(parents=True, exist_ok=True)
|
|
|
|
|
+ downloaded = snapshot_download(
|
|
|
|
|
+ "BAAI/bge-reranker-base",
|
|
|
|
|
+ cache_dir=str(_RERANKER_PATH.parent),
|
|
|
|
|
+ )
|
|
|
|
|
+ # snapshot_download 返回缓存路径,可能不是我们指定的路径
|
|
|
|
|
+ # 如果路径不同,做复制
|
|
|
|
|
+ downloaded_path = Path(downloaded)
|
|
|
|
|
+ if downloaded_path != _RERANKER_PATH:
|
|
|
|
|
+ if _RERANKER_PATH.exists():
|
|
|
|
|
+ shutil.rmtree(str(_RERANKER_PATH))
|
|
|
|
|
+ shutil.copytree(str(downloaded_path), str(_RERANKER_PATH))
|
|
|
|
|
+ print(f"[Reranker] 下载完成: {_RERANKER_PATH}")
|
|
|
|
|
+ return str(_RERANKER_PATH)
|
|
|
|
|
+ except Exception as e:
|
|
|
|
|
+ print(f"[Reranker] 下载失败: {e}")
|
|
|
|
|
+ print(f"[Reranker] 将跳过精排步骤(不影响其他功能)")
|
|
|
|
|
+ return None
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def init_reranker() -> Optional[object]:
|
|
|
|
|
+ """初始化 CrossEncoder reranker。
|
|
|
|
|
+
|
|
|
|
|
+ 返回 CrossEncoder 实例,如果不可用返回 None。
|
|
|
|
|
+ """
|
|
|
|
|
+ global _reranker
|
|
|
|
|
+ if _reranker is not None:
|
|
|
|
|
+ return _reranker if _reranker != "NONE" else None
|
|
|
|
|
+
|
|
|
|
|
+ model_path = download_reranker()
|
|
|
|
|
+ if model_path is None:
|
|
|
|
|
+ _reranker = "NONE"
|
|
|
|
|
+ return None
|
|
|
|
|
+
|
|
|
|
|
+ try:
|
|
|
|
|
+ from sentence_transformers import CrossEncoder
|
|
|
|
|
+ print(f"[Reranker] 加载模型: {model_path}")
|
|
|
|
|
+ _reranker = CrossEncoder(model_path, device="cpu")
|
|
|
|
|
+ print(f"[Reranker] 加载完成")
|
|
|
|
|
+ return _reranker
|
|
|
|
|
+ except Exception as e:
|
|
|
|
|
+ print(f"[Reranker] 加载失败: {e}")
|
|
|
|
|
+ _reranker = "NONE"
|
|
|
|
|
+ return None
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+# ============================================================================
|
|
|
|
|
+# Section 3: 文档加载 + 数据清洗 + 章节切块 (优化1, 优化4)
|
|
|
|
|
+# ============================================================================
|
|
|
|
|
+
|
|
|
|
|
+def clean_markdown_escapes(text: str) -> str:
|
|
|
|
|
+ """清理 Markdown 转义符。
|
|
|
|
|
+
|
|
|
|
|
+ Markdown 文件中常见 \-、\.、\* 等转义符,会污染关键词匹配和
|
|
|
|
|
+ Reranker 相关度计算(例如 "试用期\\-" 和 "试用期" 在 BM25 中是不同 token)。
|
|
|
|
|
+
|
|
|
|
|
+ 清理示例:
|
|
|
|
|
+ 9:00\\-18:00 → 9:00-18:00
|
|
|
|
|
+ 1\\.5小时 → 1.5小时
|
|
|
|
|
+ V2\\.0 → V2.0
|
|
|
|
|
+ """
|
|
|
|
|
+ # 去掉反斜杠转义符:匹配 \<符号> 并只保留 <符号>
|
|
|
|
|
+ # 字符类中的特殊符号需要逐个列出:- . * ? + [ ] ( ) { } ! ^ $ | #
|
|
|
|
|
+ text = re.sub(r'\\([\-.*?+\[\](){}!^$|#])', r'\1', text)
|
|
|
|
|
+ # 合并多余空行(>2 个连续换行):避免 chunks 中出现大段空白
|
|
|
|
|
+ text = re.sub(r'\n{3,}', '\n\n', text)
|
|
|
|
|
+ return text
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def load_and_chunk() -> list:
|
|
|
|
|
+ """加载员工手册 Markdown 文件,按章节结构切块。
|
|
|
|
|
+
|
|
|
|
|
+ 优化1: 使用 XX科技有限公司员工手册.md
|
|
|
|
|
+ 优化4: 按 Markdown 章节(## 第X条)切块,保证每条语义完整
|
|
|
|
|
+
|
|
|
|
|
+ 切块策略:
|
|
|
|
|
+ 1. 读取 .md 文件全部文本
|
|
|
|
|
+ 2. **数据清洗**:清理 Markdown 转义符(9:00\\-18:00 → 9:00-18:00)
|
|
|
|
|
+ 3. 跳过目录,从 # 第一章 开始
|
|
|
|
|
+ 4. 按章节标题(## 第X条)切块,每个条款作为一个独立 chunk
|
|
|
|
|
+ 5. 过长的条款二次切分(保留章节上下文前缀)
|
|
|
|
|
+
|
|
|
|
|
+ 为什么不用 SemanticChunker?
|
|
|
|
|
+ → SemanticChunker 在文档结构清晰时(如员工手册有章节)反而把相关内容
|
|
|
|
|
+ 合并过度,导致向量被稀释。按 Markdown 标题切能保证每个 chunk 语义完整。
|
|
|
|
|
+
|
|
|
|
|
+ Returns:
|
|
|
|
|
+ list[Document]: 文档块列表,每个包含 page_content 和 metadata
|
|
|
|
|
+ """
|
|
|
|
|
+ from langchain_core.documents import Document
|
|
|
|
|
+ from langchain_text_splitters import RecursiveCharacterTextSplitter
|
|
|
|
|
+
|
|
|
|
|
+ print(f"\n{'='*60}")
|
|
|
|
|
+ print(f"[文档加载] 读取: {_DATA_PATH.name}")
|
|
|
|
|
+
|
|
|
|
|
+ # Step 1: 读取文件
|
|
|
|
|
+ raw_text = _DATA_PATH.read_text(encoding="utf-8")
|
|
|
|
|
+ print(f"[文档加载] 原始长度: {len(raw_text)} 字符")
|
|
|
|
|
+
|
|
|
|
|
+ # Step 2: 数据清洗 - 清理 Markdown 转义符
|
|
|
|
|
+ cleaned_text = clean_markdown_escapes(raw_text)
|
|
|
|
|
+ print(f"[文档加载] 清洗后长度: {len(cleaned_text)} 字符")
|
|
|
|
|
+
|
|
|
|
|
+ # Step 3: 跳过目录,从正文开始
|
|
|
|
|
+ body_start = cleaned_text.find("# 第一章")
|
|
|
|
|
+ if body_start == -1:
|
|
|
|
|
+ body_start = cleaned_text.find("# 第")
|
|
|
|
|
+ if body_start > 0:
|
|
|
|
|
+ body_text = cleaned_text[body_start:]
|
|
|
|
|
+ else:
|
|
|
|
|
+ body_text = cleaned_text
|
|
|
|
|
+
|
|
|
|
|
+ # Step 4: 按章节结构切块
|
|
|
|
|
+ # 每个 "## 第X条" 作为一个独立的 chunk,保证语义完整
|
|
|
|
|
+ chunks = []
|
|
|
|
|
+ current_chapter = ""
|
|
|
|
|
+ current_section = ""
|
|
|
|
|
+ current_section_lines = []
|
|
|
|
|
+
|
|
|
|
|
+ def flush_section():
|
|
|
|
|
+ """保存当前 section 为一个 Document。"""
|
|
|
|
|
+ if not current_section or not current_section_lines:
|
|
|
|
|
+ return
|
|
|
|
|
+ content = "\n".join(current_section_lines).strip()
|
|
|
|
|
+ if len(content) < 15:
|
|
|
|
|
+ return
|
|
|
|
|
+ # 跳过纯符号行
|
|
|
|
|
+ if re.match(r"^[\s\-—#*=_.]+$", content):
|
|
|
|
|
+ return
|
|
|
|
|
+ # 附加上下文前缀:让 embedding 知道这段文字属于哪个章节
|
|
|
|
|
+ full_content = f"【{current_chapter} - {current_section}】\n{content}"
|
|
|
|
|
+ chunks.append(Document(
|
|
|
|
|
+ page_content=full_content,
|
|
|
|
|
+ metadata={
|
|
|
|
|
+ "source": _DATA_PATH.name,
|
|
|
|
|
+ "chapter": current_chapter,
|
|
|
|
|
+ "section": current_section,
|
|
|
|
|
+ }
|
|
|
|
|
+ ))
|
|
|
|
|
+
|
|
|
|
|
+ for line in body_text.split("\n"):
|
|
|
|
|
+ # 检测章标题 (# 第X章 ...) — 一级标题,更新当前所在章
|
|
|
|
|
+ chap_match = re.match(r"^#\s+(第[一二三四五六七八九十\d]+章\s*\S.*)", line)
|
|
|
|
|
+ if chap_match:
|
|
|
|
|
+ flush_section() # 先保存上一节(章切换意味着上一节结束)
|
|
|
|
|
+ current_chapter = chap_match.group(1).strip()
|
|
|
|
|
+ current_section = "" # 进入新章时清空 section,等遇到 ## 才开始记录
|
|
|
|
|
+ current_section_lines = []
|
|
|
|
|
+ continue
|
|
|
|
|
+ # 检测条标题 (## 第X条 ...) — 二级标题,每个条款是一个独立 chunk
|
|
|
|
|
+ sec_match = re.match(r"^##\s+(第[一二三四五六七八九十\d]+条\s*\S.*)", line)
|
|
|
|
|
+ if sec_match:
|
|
|
|
|
+ flush_section() # 先保存上一节(条款切换意味着上一条结束)
|
|
|
|
|
+ current_section = sec_match.group(1).strip()
|
|
|
|
|
+ current_section_lines = []
|
|
|
|
|
+ continue
|
|
|
|
|
+ # 跳过分隔线 (---):Markdown 的水平分隔符,对内容无意义
|
|
|
|
|
+ if re.match(r"^---+\s*$", line):
|
|
|
|
|
+ continue
|
|
|
|
|
+ # 收集正文:只有在已进入某个 section 后才开始收集
|
|
|
|
|
+ # (目录、文档头部说明等非条款内容会被跳过)
|
|
|
|
|
+ if current_section:
|
|
|
|
|
+ stripped = line.strip()
|
|
|
|
|
+ if stripped: # 跳过空行
|
|
|
|
|
+ current_section_lines.append(stripped)
|
|
|
|
|
+
|
|
|
|
|
+ flush_section() # 保存最后一节(循环结束后内存中还残留最后一节未写入)
|
|
|
|
|
+
|
|
|
|
|
+ print(f"[文档加载] 按章节切块完成: {len(chunks)} 个条款")
|
|
|
|
|
+
|
|
|
|
|
+ # Step 5: 对过长的 section 二次切分(保留章节上下文)
|
|
|
|
|
+ # 防止单个条款过长(>1200 字符)稀释向量
|
|
|
|
|
+ final_chunks = []
|
|
|
|
|
+ sub_splitter = RecursiveCharacterTextSplitter(
|
|
|
|
|
+ chunk_size=800,
|
|
|
|
|
+ chunk_overlap=100,
|
|
|
|
|
+ separators=["\n\n", "\n", "。", "!", "?", ";", ",", " ", ""]
|
|
|
|
|
+ )
|
|
|
|
|
+ for c in chunks:
|
|
|
|
|
+ if len(c.page_content) > 1200:
|
|
|
|
|
+ sub_chunks = sub_splitter.split_documents([c])
|
|
|
|
|
+ final_chunks.extend(sub_chunks)
|
|
|
|
|
+ else:
|
|
|
|
|
+ final_chunks.append(c)
|
|
|
|
|
+
|
|
|
|
|
+ print(f"[文档加载] 最终块数量: {len(final_chunks)} 个 "
|
|
|
|
|
+ f"(过长条款已二次切分)")
|
|
|
|
|
+
|
|
|
|
|
+ if len(final_chunks) < 10:
|
|
|
|
|
+ print(f"[文档加载] 警告: 块数量过少 ({len(final_chunks)})")
|
|
|
|
|
+
|
|
|
|
|
+ return final_chunks
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+# ============================================================================
|
|
|
|
|
+# Section 4: ChromaDB 向量存储 (优化3)
|
|
|
|
|
+# ============================================================================
|
|
|
|
|
+
|
|
|
|
|
+def build_vectorstore(chunks: list) -> object:
|
|
|
|
|
+ """构建 ChromaDB 向量存储。
|
|
|
|
|
+
|
|
|
|
|
+ 使用 BGE-Large-Zh-v1.5 embedding 模型 (1024d),collection 配置余弦相似度。
|
|
|
|
|
+ 每次调用强制重建(开发阶段)。
|
|
|
|
|
+ """
|
|
|
|
|
+ from langchain_community.vectorstores import Chroma
|
|
|
|
|
+
|
|
|
|
|
+ print(f"\n[向量存储] 构建 ChromaDB ({len(chunks)} 个文档)")
|
|
|
|
|
+ embedding_model = init_embedding_model()
|
|
|
|
|
+
|
|
|
|
|
+ # 清理旧数据
|
|
|
|
|
+ if _CHROMA_DIR.exists():
|
|
|
|
|
+ print(f"[向量存储] 清理旧数据库: {_CHROMA_DIR}")
|
|
|
|
|
+ shutil.rmtree(str(_CHROMA_DIR))
|
|
|
|
|
+
|
|
|
|
|
+ vectorstore = Chroma.from_documents(
|
|
|
|
|
+ documents=chunks,
|
|
|
|
|
+ embedding=embedding_model,
|
|
|
|
|
+ collection_metadata={"hnsw:space": "cosine"},
|
|
|
|
|
+ persist_directory=str(_CHROMA_DIR),
|
|
|
|
|
+ )
|
|
|
|
|
+ print(f"[向量存储] 构建完成,持久化到: {_CHROMA_DIR}")
|
|
|
|
|
+ return vectorstore
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def get_vectorstore(chunks: list) -> object:
|
|
|
|
|
+ """获取向量存储(优先加载已有,否则重建)。"""
|
|
|
|
|
+ from langchain_community.vectorstores import Chroma
|
|
|
|
|
+
|
|
|
|
|
+ sqlite_file = _CHROMA_DIR / "chroma.sqlite3"
|
|
|
|
|
+ if sqlite_file.exists():
|
|
|
|
|
+ embedding_model = init_embedding_model()
|
|
|
|
|
+ try:
|
|
|
|
|
+ print(f"[向量存储] 从 {_CHROMA_DIR} 加载已有数据库")
|
|
|
|
|
+ return Chroma(
|
|
|
|
|
+ persist_directory=str(_CHROMA_DIR),
|
|
|
|
|
+ embedding_function=embedding_model,
|
|
|
|
|
+ collection_metadata={"hnsw:space": "cosine"},
|
|
|
|
|
+ )
|
|
|
|
|
+ except Exception as e:
|
|
|
|
|
+ print(f"[向量存储] 加载失败 ({e}),将重建")
|
|
|
|
|
+
|
|
|
|
|
+ return build_vectorstore(chunks)
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+# ============================================================================
|
|
|
|
|
+# Section 5: 多路召回系统 (优化5)
|
|
|
|
|
+# ============================================================================
|
|
|
|
|
+
|
|
|
|
|
+_bm25_retriever = None
|
|
|
|
|
+_vector_retriever = None
|
|
|
|
|
+_ensemble_retriever = None
|
|
|
|
|
+_mq_retriever = None
|
|
|
|
|
+
|
|
|
|
|
+def _jieba_tokenize(text: str) -> list[str]:
|
|
|
|
|
+ """使用 jieba 对中文文本分词。
|
|
|
|
|
+
|
|
|
|
|
+ BM25Retriever 默认用英文空白分词器,遇到中文时会将整段文本当作一个 token,
|
|
|
|
|
+ 导致完全无法匹配关键词(不同查询返回相同结果)。
|
|
|
|
|
+ 用 jieba 分词后,"试用期是几个月" 会被切成 ["试用期", "是", "几个", "月"],
|
|
|
|
|
+ BM25 才能正常按词频打分。
|
|
|
|
|
+ """
|
|
|
|
|
+ import jieba
|
|
|
|
|
+ # lcut 精确模式(vs cut 全模式):返回 list,每个 token 是一个词
|
|
|
|
|
+ tokens = jieba.lcut(text)
|
|
|
|
|
+ # 过滤规则:去空白 + 丢弃单字
|
|
|
|
|
+ # 单字(如"的""是""一")信息量极低,且大量重复,会拉低 BM25 区分度
|
|
|
|
|
+ return [t for t in tokens if t.strip() and len(t.strip()) > 1]
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def create_bm25_retriever(chunks: list):
|
|
|
|
|
+ """创建 BM25 关键词检索器。
|
|
|
|
|
+
|
|
|
|
|
+ BM25 基于词频统计,擅长精确关键词匹配,
|
|
|
|
|
+ 弥补向量检索在专有名词、数字、缩写上的短板。
|
|
|
|
|
+
|
|
|
|
|
+ 注意:必须传入中文分词器 _jieba_tokenize,否则默认英文分词器
|
|
|
|
|
+ 会把整段中文当成一个 token,BM25 退化为按原始顺序返回。
|
|
|
|
|
+ """
|
|
|
|
|
+ global _bm25_retriever
|
|
|
|
|
+ if _bm25_retriever is not None:
|
|
|
|
|
+ return _bm25_retriever
|
|
|
|
|
+ from langchain_community.retrievers import BM25Retriever
|
|
|
|
|
+ _bm25_retriever = BM25Retriever.from_documents(
|
|
|
|
|
+ chunks,
|
|
|
|
|
+ preprocess_func=_jieba_tokenize,
|
|
|
|
|
+ )
|
|
|
|
|
+ _bm25_retriever.k = RETRIEVAL_K
|
|
|
|
|
+ print(f"[检索器] BM25 初始化完成 (k={RETRIEVAL_K}, 中文分词=jieba)")
|
|
|
|
|
+ return _bm25_retriever
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def create_vector_retriever(vectorstore):
|
|
|
|
|
+ """创建向量语义检索器。
|
|
|
|
|
+
|
|
|
|
|
+ 基于 BGE-Large-Zh-v1.5 embedding 的余弦相似度检索,
|
|
|
|
|
+ 能在语义层面理解同义词、近义词和概念关联。
|
|
|
|
|
+ """
|
|
|
|
|
+ global _vector_retriever
|
|
|
|
|
+ if _vector_retriever is not None:
|
|
|
|
|
+ return _vector_retriever
|
|
|
|
|
+ _vector_retriever = vectorstore.as_retriever(
|
|
|
|
|
+ search_type="similarity",
|
|
|
|
|
+ search_kwargs={"k": RETRIEVAL_K},
|
|
|
|
|
+ )
|
|
|
|
|
+ print(f"[检索器] 向量检索器初始化完成 (k={RETRIEVAL_K})")
|
|
|
|
|
+ return _vector_retriever
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def create_ensemble_retriever(chunks, vectorstore):
|
|
|
|
|
+ """创建混合检索器: BM25 + 向量检索融合。
|
|
|
|
|
+
|
|
|
|
|
+ - BM25 权重 0.3: 关键词精确匹配
|
|
|
|
|
+ - 向量权重 0.7: 语义理解(权重更高,因为语义匹配通常更重要)
|
|
|
|
|
+
|
|
|
|
|
+ 融合方式: 分数归一化后按权重加权求和。
|
|
|
|
|
+ """
|
|
|
|
|
+ global _ensemble_retriever
|
|
|
|
|
+ if _ensemble_retriever is not None:
|
|
|
|
|
+ return _ensemble_retriever
|
|
|
|
|
+ from langchain_classic.retrievers import EnsembleRetriever
|
|
|
|
|
+
|
|
|
|
|
+ bm25 = create_bm25_retriever(chunks)
|
|
|
|
|
+ vec = create_vector_retriever(vectorstore)
|
|
|
|
|
+ _ensemble_retriever = EnsembleRetriever(
|
|
|
|
|
+ retrievers=[bm25, vec],
|
|
|
|
|
+ weights=ENSEMBLE_WEIGHTS,
|
|
|
|
|
+ )
|
|
|
|
|
+ print(f"[检索器] 混合检索器初始化完成 "
|
|
|
|
|
+ f"(BM25:{ENSEMBLE_WEIGHTS[0]} + 向量:{ENSEMBLE_WEIGHTS[1]})")
|
|
|
|
|
+ return _ensemble_retriever
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def create_mq_retriever(vectorstore):
|
|
|
|
|
+ """创建多查询检索器 (Multi-Query Retriever)。
|
|
|
|
|
+
|
|
|
|
|
+ LLM 自动生成 3 个查询变体 → 分别检索 → 合并去重。
|
|
|
|
|
+ 适合处理用户问题表述不够精确的场景。
|
|
|
|
|
+ """
|
|
|
|
|
+ global _mq_retriever
|
|
|
|
|
+ if _mq_retriever is not None:
|
|
|
|
|
+ return _mq_retriever
|
|
|
|
|
+ from langchain_classic.retrievers import MultiQueryRetriever
|
|
|
|
|
+
|
|
|
|
|
+ vec = create_vector_retriever(vectorstore)
|
|
|
|
|
+ _mq_retriever = MultiQueryRetriever.from_llm(
|
|
|
|
|
+ retriever=vec,
|
|
|
|
|
+ llm=get_llm(),
|
|
|
|
|
+ )
|
|
|
|
|
+ print(f"[检索器] 多查询检索器初始化完成 (变体数={MULTI_QUERY_VARIANTS})")
|
|
|
|
|
+ return _mq_retriever
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def retrieve(query: str, strategy: str = "ensemble",
|
|
|
|
|
+ top_k: int = DEFAULT_K) -> list:
|
|
|
|
|
+ """统一的检索调度接口。
|
|
|
|
|
+
|
|
|
|
|
+ Args:
|
|
|
|
|
+ query: 用户问题
|
|
|
|
|
+ strategy: 检索策略,可选 "bm25" | "vector" | "ensemble" | "multi_query"
|
|
|
|
|
+ top_k: 返回的文档数量
|
|
|
|
|
+
|
|
|
|
|
+ Returns:
|
|
|
|
|
+ list[Document]: 检索到的文档列表
|
|
|
|
|
+ """
|
|
|
|
|
+ vs = _get_vectorstore_ref()
|
|
|
|
|
+ chunks = _get_chunks_ref()
|
|
|
|
|
+
|
|
|
|
|
+ if strategy == "bm25":
|
|
|
|
|
+ ret = create_bm25_retriever(chunks)
|
|
|
|
|
+ elif strategy == "vector":
|
|
|
|
|
+ ret = create_vector_retriever(vs)
|
|
|
|
|
+ elif strategy == "multi_query":
|
|
|
|
|
+ ret = create_mq_retriever(vs)
|
|
|
|
|
+ else:
|
|
|
|
|
+ ret = create_ensemble_retriever(chunks, vs)
|
|
|
|
|
+
|
|
|
|
|
+ # 为不同检索器适配 k 参数
|
|
|
|
|
+ # BM25Retriever / VectorStoreRetriever 用 .k 属性控制返回数量
|
|
|
|
|
+ # MultiQueryRetriever / EnsembleRetriever 没有可变 .k,需要 invoke 后切片
|
|
|
|
|
+ if hasattr(ret, "k"):
|
|
|
|
|
+ # 临时修改 k → 调用 → 恢复原值(防止影响后续其他调用)
|
|
|
|
|
+ old_k = ret.k
|
|
|
|
|
+ ret.k = top_k
|
|
|
|
|
+ docs = ret.invoke(query)
|
|
|
|
|
+ ret.k = old_k
|
|
|
|
|
+ else:
|
|
|
|
|
+ docs = ret.invoke(query)
|
|
|
|
|
+ docs = docs[:top_k]
|
|
|
|
|
+
|
|
|
|
|
+ return docs
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+# 模块级引用(由主流程设置)
|
|
|
|
|
+_VS = None
|
|
|
|
|
+_CHUNKS = None
|
|
|
|
|
+
|
|
|
|
|
+def _get_vectorstore_ref():
|
|
|
|
|
+ if _VS is None:
|
|
|
|
|
+ raise RuntimeError("向量存储未初始化,请先调用 load_and_chunk() 和 get_vectorstore()")
|
|
|
|
|
+ return _VS
|
|
|
|
|
+
|
|
|
|
|
+def _get_chunks_ref():
|
|
|
|
|
+ if _CHUNKS is None:
|
|
|
|
|
+ raise RuntimeError("文档块未初始化")
|
|
|
|
|
+ return _CHUNKS
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+# ============================================================================
|
|
|
|
|
+# Section 6: 查询增强 (优化6)
|
|
|
|
|
+# ============================================================================
|
|
|
|
|
+
|
|
|
|
|
+def rewrite_query(original: str) -> str:
|
|
|
|
|
+ """查询重写:用 LLM 将用户问题改写为更精确的检索查询。
|
|
|
|
|
+
|
|
|
|
|
+ 改写策略:
|
|
|
|
|
+ - 提取核心实体和关键词
|
|
|
|
|
+ - 消除歧义,明确指代
|
|
|
|
|
+ - 将口语化表达转为文档化风格
|
|
|
|
|
+ """
|
|
|
|
|
+ llm = get_llm()
|
|
|
|
|
+ prompt = f"""你是一个查询优化助手。给定用户的原始问题,请将其改写成更适合文档检索的形式。
|
|
|
|
|
+
|
|
|
|
|
+要求:
|
|
|
|
|
+1. 提取核心实体和关键词
|
|
|
|
|
+2. 移除歧义,明确指代
|
|
|
|
|
+3. 保持原意不变
|
|
|
|
|
+4. 使用文档中可能出现的术语(如"年假""加班费""绩效考核"等正式表述)
|
|
|
|
|
+5. 直接输出改写后的查询,不要任何解释
|
|
|
|
|
+
|
|
|
|
|
+原始问题:{original}
|
|
|
|
|
+
|
|
|
|
|
+改写后的查询:"""
|
|
|
|
|
+
|
|
|
|
|
+ try:
|
|
|
|
|
+ raw = llm.invoke(prompt)
|
|
|
|
|
+ rewritten = raw.content if hasattr(raw, "content") else str(raw)
|
|
|
|
|
+ rewritten = rewritten.strip().strip('"').strip("'")
|
|
|
|
|
+ if len(rewritten) < 2:
|
|
|
|
|
+ return original
|
|
|
|
|
+ return rewritten
|
|
|
|
|
+ except Exception as e:
|
|
|
|
|
+ print(f"[查询重写] 失败 ({e}),使用原始查询")
|
|
|
|
|
+ return original
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def hyde_retrieve(query: str, vectorstore, top_k: int = RETRIEVAL_K) -> list:
|
|
|
|
|
+ """HyDE (Hypothetical Document Embeddings) 检索。
|
|
|
|
|
+
|
|
|
|
|
+ 核心思路: LLM 先生成一个「假设性答案」→ 将这个答案向量化 →
|
|
|
|
|
+ 用这个向量去检索 → 假设性答案的语义空间比问题本身更接近真实文档。
|
|
|
|
|
+
|
|
|
|
|
+ 流程:
|
|
|
|
|
+ 1. LLM 生成假设性答案(不需要准确,只需"看起来像"文档内容)
|
|
|
|
|
+ 2. Embedding 模型将假设性答案转为向量
|
|
|
|
|
+ 3. 在 ChromaDB 中检索相似文档
|
|
|
|
|
+ """
|
|
|
|
|
+ llm = get_llm()
|
|
|
|
|
+ prompt = f"""请根据以下问题,生成一个假设性的员工手册回答。这个回答不需要完全准确,
|
|
|
|
|
+只需要看起来像是一段真实的公司制度文档内容。用来帮助搜索引擎找到相关文档。
|
|
|
|
|
+
|
|
|
|
|
+要求:直接输出内容,不要任何前缀或解释。
|
|
|
|
|
+
|
|
|
|
|
+问题:{query}
|
|
|
|
|
+
|
|
|
|
|
+假设文档:"""
|
|
|
|
|
+
|
|
|
|
|
+ try:
|
|
|
|
|
+ raw = llm.invoke(prompt)
|
|
|
|
|
+ hypothesis = raw.content if hasattr(raw, "content") else str(raw)
|
|
|
|
|
+ hypothesis = hypothesis.strip()
|
|
|
|
|
+ if len(hypothesis) < 10:
|
|
|
|
|
+ return []
|
|
|
|
|
+
|
|
|
|
|
+ docs = vectorstore.similarity_search(hypothesis, k=top_k)
|
|
|
|
|
+ return docs
|
|
|
|
|
+ except Exception as e:
|
|
|
|
|
+ print(f"[HyDE] 失败 ({e})")
|
|
|
|
|
+ return []
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def reciprocal_rank_fusion(doc_lists: list[list], k: int = RRF_K) -> list:
|
|
|
|
|
+ """倒数秩融合 (Reciprocal Rank Fusion)。
|
|
|
|
|
+
|
|
|
|
|
+ 将多个检索结果列表按 RRF 算法合并为一个去重排序列表。
|
|
|
|
|
+ 公式: score(d) = Σ 1/(rank_i(d) + k)
|
|
|
|
|
+
|
|
|
|
|
+ Args:
|
|
|
|
|
+ doc_lists: 多个检索结果列表
|
|
|
|
|
+ k: 平滑参数(默认60,标准设置)
|
|
|
|
|
+
|
|
|
|
|
+ Returns:
|
|
|
|
|
+ 融合后按 RRF 分数降序排列的文档列表
|
|
|
|
|
+ """
|
|
|
|
|
+ scores: dict[int, float] = {}
|
|
|
|
|
+ doc_map: dict[int, object] = {}
|
|
|
|
|
+
|
|
|
|
|
+ # Step 1: 遍历每个检索器返回的列表,按排名累计 RRF 分数
|
|
|
|
|
+ for doc_list in doc_lists:
|
|
|
|
|
+ for rank, doc in enumerate(doc_list): # rank 从 0 开始,表示该文档在该列表中的位置
|
|
|
|
|
+ # 用 page_content 的 hash 作为文档唯一标识(同一段文字会得到相同 hash)
|
|
|
|
|
+ doc_id = hash(doc.page_content)
|
|
|
|
|
+ if doc_id not in scores:
|
|
|
|
|
+ scores[doc_id] = 0.0
|
|
|
|
|
+ doc_map[doc_id] = doc
|
|
|
|
|
+ # RRF 核心公式:score(d) = Σ 1/(rank + k)
|
|
|
|
|
+ # rank 从 0 开始,所以要 +1 才是真正的「第几名」
|
|
|
|
|
+ # k 是平滑常数:避免 rank=0 的文档得到 1.0 满分,让排名靠后的也有贡献
|
|
|
|
|
+ # k=60 是业界标准值,源自论文《Reciprocal Rank Fusion outperforms Condorcet》
|
|
|
|
|
+ scores[doc_id] += 1.0 / (rank + k)
|
|
|
|
|
+
|
|
|
|
|
+ # Step 2: 按 RRF 分数降序排列,输出融合后的文档列表
|
|
|
|
|
+ sorted_ids = sorted(scores.keys(), key=lambda x: scores[x], reverse=True)
|
|
|
|
|
+ return [doc_map[did] for did in sorted_ids]
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def enhanced_retrieve(query: str, strategy: str = "ensemble",
|
|
|
|
|
+ top_k: int = DEFAULT_K) -> list:
|
|
|
|
|
+ """增强检索: 查询重写 + HyDE + RRF 融合。
|
|
|
|
|
+
|
|
|
|
|
+ 流程:
|
|
|
|
|
+ 1. LLM 重写查询 → 得到更精确的检索查询
|
|
|
|
|
+ 2. HyDE 生成假设答案 → 从另一个语义角度检索
|
|
|
|
|
+ 3. 标准检索(用重写后的查询)
|
|
|
|
|
+ 4. RRF 融合 HyDE 结果 + 标准结果 → 取 top_k
|
|
|
|
|
+ """
|
|
|
|
|
+ vectorstore = _get_vectorstore_ref()
|
|
|
|
|
+
|
|
|
|
|
+ # Step 1: 查询重写
|
|
|
|
|
+ rewritten = rewrite_query(query)
|
|
|
|
|
+ if rewritten != query:
|
|
|
|
|
+ print(f"[增强检索] 查询重写: 「{query[:40]}...」→「{rewritten[:60]}...」")
|
|
|
|
|
+
|
|
|
|
|
+ # Step 2: HyDE 检索
|
|
|
|
|
+ hyde_docs = hyde_retrieve(query, vectorstore, top_k=RETRIEVAL_K)
|
|
|
|
|
+ print(f"[增强检索] HyDE 召回 {len(hyde_docs)} 个文档")
|
|
|
|
|
+
|
|
|
|
|
+ # Step 3: 标准检索(使用重写后的查询)
|
|
|
|
|
+ standard_docs = retrieve(rewritten, strategy=strategy, top_k=RETRIEVAL_K)
|
|
|
|
|
+ print(f"[增强检索] {strategy} 召回 {len(standard_docs)} 个文档")
|
|
|
|
|
+
|
|
|
|
|
+ # Step 4: RRF 融合
|
|
|
|
|
+ if hyde_docs and standard_docs:
|
|
|
|
|
+ fused = reciprocal_rank_fusion([standard_docs, hyde_docs])
|
|
|
|
|
+ print(f"[增强检索] RRF 融合后 {len(fused)} 个文档")
|
|
|
|
|
+ elif standard_docs:
|
|
|
|
|
+ fused = standard_docs
|
|
|
|
|
+ else:
|
|
|
|
|
+ fused = hyde_docs
|
|
|
|
|
+
|
|
|
|
|
+ return fused[:top_k]
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+# ============================================================================
|
|
|
|
|
+# Section 7: Reranker 精排 (优化7)
|
|
|
|
|
+# ============================================================================
|
|
|
|
|
+
|
|
|
|
|
+def rerank(query: str, candidates: list, top_k: int = DEFAULT_K) -> list:
|
|
|
|
|
+ """使用 BGE-Reranker 对候选文档精排。
|
|
|
|
|
+
|
|
|
|
|
+ Cross-Encoder 将 query-doc 对联合编码,比 embedding 相似度更精确。
|
|
|
|
|
+ 缺点: 速度较慢,适合对 Top-N 候选做精排(N 通常 10-20)。
|
|
|
|
|
+
|
|
|
|
|
+ Args:
|
|
|
|
|
+ query: 用户问题
|
|
|
|
|
+ candidates: 候选文档列表
|
|
|
|
|
+ top_k: 返回的文档数量
|
|
|
|
|
+
|
|
|
|
|
+ Returns:
|
|
|
|
|
+ 按相关度分数降序排列的文档列表
|
|
|
|
|
+ """
|
|
|
|
|
+ reranker = init_reranker()
|
|
|
|
|
+ if reranker is None or len(candidates) == 0:
|
|
|
|
|
+ return candidates[:top_k]
|
|
|
|
|
+
|
|
|
|
|
+ # 优化:候选数 ≤ top_k 时无需精排,直接返回(省下推理时间)
|
|
|
|
|
+ if len(candidates) <= top_k:
|
|
|
|
|
+ return candidates
|
|
|
|
|
+
|
|
|
|
|
+ try:
|
|
|
|
|
+ # 构建 query-doc 配对:CrossEncoder 是双塔模型的对比
|
|
|
|
|
+ # 与 embedding(Bi-Encoder,query 和 doc 分别编码再算相似度)不同,
|
|
|
|
|
+ # CrossEncoder 把 query 和 doc 拼在一起送入模型,深度交互注意力,
|
|
|
|
|
+ # 因此精度更高但速度更慢(适合 Top-N 精排,N 通常 10-20)
|
|
|
|
|
+ pairs = [(query, doc.page_content) for doc in candidates]
|
|
|
|
|
+ scores = reranker.predict(pairs, show_progress_bar=False)
|
|
|
|
|
+
|
|
|
|
|
+ # 按相关度分数降序排列,取 top_k
|
|
|
|
|
+ # scored 是 [(score, doc), ...],sort 后取前 top_k 个的 doc
|
|
|
|
|
+ scored = list(zip(scores, candidates))
|
|
|
|
|
+ scored.sort(key=lambda x: x[0], reverse=True)
|
|
|
|
|
+
|
|
|
|
|
+ reranked = [doc for _, doc in scored[:top_k]]
|
|
|
|
|
+ print(f"[Reranker] 精排完成: {len(candidates)} → {len(reranked)} "
|
|
|
|
|
+ f"(最高分: {scored[0][0]:.4f})")
|
|
|
|
|
+ return reranked
|
|
|
|
|
+ except Exception as e:
|
|
|
|
|
+ print(f"[Reranker] 精排失败 ({e}),使用原始顺序")
|
|
|
|
|
+ return candidates[:top_k]
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+# ============================================================================
|
|
|
|
|
+# Section 8: Agent 工具封装 (优化2)
|
|
|
|
|
+# ============================================================================
|
|
|
|
|
+
|
|
|
|
|
+def create_retrieval_tool():
|
|
|
|
|
+ """将检索系统封装为 LangChain Tool。
|
|
|
|
|
+
|
|
|
|
|
+ Agent 可根据问题自动判断是否需要调用此工具检索员工手册。
|
|
|
|
|
+ """
|
|
|
|
|
+ from langchain_core.tools import tool
|
|
|
|
|
+
|
|
|
|
|
+ @tool
|
|
|
|
|
+ def search_employee_handbook(query: str) -> str:
|
|
|
|
|
+ """从《XX科技有限公司员工手册》中检索与 query 相关的内容。
|
|
|
|
|
+
|
|
|
|
|
+ 当用户问到以下内容时使用此工具:
|
|
|
|
|
+ - 公司制度、考勤规定、请假流程
|
|
|
|
|
+ - 薪资福利、社保公积金
|
|
|
|
|
+ - 加班管理、出差报销
|
|
|
|
|
+ - 绩效考核、晋升条件
|
|
|
|
|
+ - 离职手续、保密规定
|
|
|
|
|
+ - 员工关怀、培训发展
|
|
|
|
|
+ - 入职流程、试用期规定
|
|
|
|
|
+
|
|
|
|
|
+ 输入:用户的问题或关键词。
|
|
|
|
|
+ 输出:手册中的相关章节内容。
|
|
|
|
|
+ """
|
|
|
|
|
+ # 使用增强检索 + 精排
|
|
|
|
|
+ docs = enhanced_retrieve(query, strategy="ensemble", top_k=RETRIEVAL_K)
|
|
|
|
|
+ docs = rerank(query, docs, top_k=DEFAULT_K)
|
|
|
|
|
+
|
|
|
|
|
+ parts = []
|
|
|
|
|
+ for i, doc in enumerate(docs):
|
|
|
|
|
+ parts.append(f"【参考条款 {i+1}】\n{doc.page_content}")
|
|
|
|
|
+ return "\n\n---\n\n".join(parts)
|
|
|
|
|
+
|
|
|
|
|
+ return search_employee_handbook
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def create_rag_agent():
|
|
|
|
|
+ """创建 RAG Agent。
|
|
|
|
|
+
|
|
|
|
|
+ Agent 能自动判断:
|
|
|
|
|
+ - 需要查手册时 → 调用 search_employee_handbook 工具
|
|
|
|
|
+ - 闲聊/通用问题 → 直接回答
|
|
|
|
|
+ """
|
|
|
|
|
+ from langchain.agents import create_agent
|
|
|
|
|
+
|
|
|
|
|
+ tool = create_retrieval_tool()
|
|
|
|
|
+ llm = get_llm()
|
|
|
|
|
+ agent = create_agent(
|
|
|
|
|
+ model=llm,
|
|
|
|
|
+ tools=[tool],
|
|
|
|
|
+ system_prompt=(
|
|
|
|
|
+ "你是XX科技有限公司的员工手册问答助手。\n\n"
|
|
|
|
|
+ "规则:\n"
|
|
|
|
|
+ "1. 当用户问到公司制度、考勤、请假、薪资、福利、加班、离职、保密等"
|
|
|
|
|
+ "规定时,必须先调用 search_employee_handbook 工具检索手册内容。\n"
|
|
|
|
|
+ "2. 基于检索到的内容回答,不要编造。如果手册中没有相关信息,"
|
|
|
|
|
+ "如实告知用户。\n"
|
|
|
|
|
+ "3. 回答要简洁清晰,引用手册中的具体条款编号。\n"
|
|
|
|
|
+ "4. 对于闲聊或与员工手册无关的问题,直接回答即可。"
|
|
|
|
|
+ ),
|
|
|
|
|
+ )
|
|
|
|
|
+ print(f"[Agent] RAG Agent 创建完成")
|
|
|
|
|
+ return agent
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+# ============================================================================
|
|
|
|
|
+# Section 9: 完整 RAG 链路
|
|
|
|
|
+# ============================================================================
|
|
|
|
|
+
|
|
|
|
|
+def build_context(docs: list) -> str:
|
|
|
|
|
+ """将检索到的文档拼接为 LLM 上下文。"""
|
|
|
|
|
+ parts = []
|
|
|
|
|
+ for i, doc in enumerate(docs):
|
|
|
|
|
+ parts.append(f"[参考文档 {i+1}]\n{doc.page_content}")
|
|
|
|
|
+ return "\n\n---\n\n".join(parts)
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def rag_query(query: str, strategy: str = "ensemble",
|
|
|
|
|
+ use_enhance: bool = True, use_rerank: bool = True) -> dict:
|
|
|
|
|
+ """确定性 RAG 查询流水线(不使用 Agent)。
|
|
|
|
|
+
|
|
|
|
|
+ Args:
|
|
|
|
|
+ query: 用户问题
|
|
|
|
|
+ strategy: 检索策略
|
|
|
|
|
+ use_enhance: 是否使用查询增强(重写 + HyDE)
|
|
|
|
|
+ use_rerank: 是否使用 Reranker 精排
|
|
|
|
|
+
|
|
|
|
|
+ Returns:
|
|
|
|
|
+ dict: {"answer": str, "sources": list[Document], "strategy": str}
|
|
|
|
|
+ """
|
|
|
|
|
+ from langchain_core.prompts import ChatPromptTemplate
|
|
|
|
|
+ from langchain_core.output_parsers import StrOutputParser
|
|
|
|
|
+
|
|
|
|
|
+ # Step 1: 检索(增强检索走 enhanced_retrieve,普通检索走 retrieve)
|
|
|
|
|
+ if use_enhance:
|
|
|
|
|
+ docs = enhanced_retrieve(query, strategy=strategy, top_k=RETRIEVAL_K)
|
|
|
|
|
+ else:
|
|
|
|
|
+ docs = retrieve(query, strategy=strategy, top_k=RETRIEVAL_K)
|
|
|
|
|
+
|
|
|
|
|
+ # Step 2: 精排(可选):对召回的 top-K 候选重新打分排序,取 top-k
|
|
|
|
|
+ if use_rerank:
|
|
|
|
|
+ docs = rerank(query, docs, top_k=DEFAULT_K)
|
|
|
|
|
+
|
|
|
|
|
+ # Step 3: 组装 prompt
|
|
|
|
|
+ # ChatPromptTemplate.from_messages: 支持 system + user 多轮对话格式
|
|
|
|
|
+ # {context} 和 {query} 是占位符,invoke 时填充
|
|
|
|
|
+ context = build_context(docs)
|
|
|
|
|
+ prompt = ChatPromptTemplate.from_messages([
|
|
|
|
|
+ ("system", """你是XX科技有限公司的员工手册问答助手。请根据以下参考资料回答问题。
|
|
|
|
|
+
|
|
|
|
|
+**规则:**
|
|
|
|
|
+- 只基于提供的参考资料回答,不要编造
|
|
|
|
|
+- 如果参考资料中没有相关信息,直接说「根据现有员工手册,我找不到这个问题的答案」
|
|
|
|
|
+- 回答要简洁清晰,引用具体条款时注明出处
|
|
|
|
|
+- 如果参考资料不足以回答,可以说明需要查阅手册的哪个章节
|
|
|
|
|
+
|
|
|
|
|
+**参考资料:**
|
|
|
|
|
+{context}"""),
|
|
|
|
|
+ ("user", "{query}")
|
|
|
|
|
+ ])
|
|
|
|
|
+
|
|
|
|
|
+ # Step 4: 生成
|
|
|
|
|
+ # LCEL 链式语法:prompt | llm | parser
|
|
|
|
|
+ # - prompt.invoke({...}) → 填充占位符生成最终 prompt
|
|
|
|
|
+ # - llm.invoke(prompt) → 调用 LLM 生成回复
|
|
|
|
|
+ # - StrOutputParser().invoke(resp) → 从 AIMessage 中提取纯文本
|
|
|
|
|
+ # 等价于:parser.invoke(llm.invoke(prompt.invoke({...})))
|
|
|
|
|
+ llm = get_llm()
|
|
|
|
|
+ chain = prompt | llm | StrOutputParser()
|
|
|
|
|
+ answer = chain.invoke({"context": context, "query": query})
|
|
|
|
|
+
|
|
|
|
|
+ return {"answer": answer, "sources": docs, "strategy": strategy}
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def agent_query(query: str) -> str:
|
|
|
|
|
+ """Agent 模式查询(自动判断是否需要检索)。
|
|
|
|
|
+
|
|
|
|
|
+ Agent 根据问题内容自动决定:
|
|
|
|
|
+ - 调用 search_employee_handbook 工具检索
|
|
|
|
|
+ - 直接回答(不检索)
|
|
|
|
|
+ """
|
|
|
|
|
+ agent = create_rag_agent()
|
|
|
|
|
+ result = agent.invoke({"messages": [("user", query)]})
|
|
|
|
|
+ # 提取最后一条消息
|
|
|
|
|
+ messages = result.get("messages", [])
|
|
|
|
|
+ if messages:
|
|
|
|
|
+ return messages[-1].content
|
|
|
|
|
+ return "Agent 未返回结果"
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+# ============================================================================
|
|
|
|
|
+# Section 10: RAG 效果评估体系 (优化8)
|
|
|
|
|
+# ============================================================================
|
|
|
|
|
+#
|
|
|
|
|
+# 评估策略说明:
|
|
|
|
|
+# 1. 手动标注测试集: 覆盖员工手册全部 16 章的真实问题 → Section 级检索评估
|
|
|
|
|
+# 2. 检索交叉验证: LLM 逐条判断检索到的文档能否回答该问题
|
|
|
|
|
+# 3. LLM 辅助生成评估 + Bad Case 根因分析
|
|
|
|
|
+
|
|
|
|
|
+# ========================================================================
|
|
|
|
|
+# 10a. 手动标注真实测试集(覆盖全部 16 章,~30% 文档覆盖率)
|
|
|
|
|
+# ========================================================================
|
|
|
|
|
+# 每条测试用例包含:
|
|
|
|
|
+# question: 员工真实会问的问题
|
|
|
|
|
+# section_keywords: 期望匹配的章节关键词(用于 Section 级检索评估)
|
|
|
|
|
+# reference_answer: 从手册原文提取的参考答案(用于生成评估)
|
|
|
|
|
+
|
|
|
|
|
+_MANUAL_TEST_SET = [
|
|
|
|
|
+ # === 第一章 总则 ===
|
|
|
|
|
+ {
|
|
|
|
|
+ "question": "员工手册制定的法律依据是什么?",
|
|
|
|
|
+ "section_keywords": ["总则", "劳动法", "劳动合同法"],
|
|
|
|
|
+ "reference_answer": "根据《中华人民共和国劳动法》《中华人民共和国劳动合同法》及相关法律法规制定。",
|
|
|
|
|
+ },
|
|
|
|
|
+ # === 第二章 入职与试用期 ===
|
|
|
|
|
+ {
|
|
|
|
|
+ "question": "试用期是多长时间?",
|
|
|
|
|
+ "section_keywords": ["试用期", "入职"],
|
|
|
|
|
+ "reference_answer": "新员工试用期为3个月,试用期包含在劳动合同期限内。",
|
|
|
|
|
+ },
|
|
|
|
|
+ {
|
|
|
|
|
+ "question": "入职需要准备哪些材料?",
|
|
|
|
|
+ "section_keywords": ["入职", "入职流程", "报到"],
|
|
|
|
|
+ "reference_answer": "入职需准备:身份证、学历学位证书、离职证明、体检报告、银行卡。",
|
|
|
|
|
+ },
|
|
|
|
|
+ # === 第三章 工作时间与考勤 ===
|
|
|
|
|
+ {
|
|
|
|
|
+ "question": "公司标准工作时间是怎样的?",
|
|
|
|
|
+ "section_keywords": ["工作时间", "考勤", "上下班"],
|
|
|
|
|
+ "reference_answer": "标准工作时间为周一至周五9:00-18:00,午休1小时,每天工作8小时。",
|
|
|
|
|
+ },
|
|
|
|
|
+ # === 第四章 假期管理 ===
|
|
|
|
|
+ {
|
|
|
|
|
+ "question": "年假没休完可以结转到下一年吗?",
|
|
|
|
|
+ "section_keywords": ["年假", "带薪年假", "假期"],
|
|
|
|
|
+ "reference_answer": "当年未休完的年假最多可结转5天至次年3月31日,逾期作废。",
|
|
|
|
|
+ },
|
|
|
|
|
+ {
|
|
|
|
|
+ "question": "请病假需要提供什么材料?",
|
|
|
|
|
+ "section_keywords": ["病假", "请假", "假期管理"],
|
|
|
|
|
+ "reference_answer": "请病假需提供二级甲等以上医院出具的病假证明和诊断证明。",
|
|
|
|
|
+ },
|
|
|
|
|
+ # === 第五章 薪资福利 ===
|
|
|
|
|
+ {
|
|
|
|
|
+ "question": "社保养老保险个人缴纳比例是多少?",
|
|
|
|
|
+ "section_keywords": ["社保", "社会保险", "五险一金", "薪资福利"],
|
|
|
|
|
+ "reference_answer": "养老保险个人缴纳比例为8%,单位缴纳比例为16%。",
|
|
|
|
|
+ },
|
|
|
|
|
+ # === 第六章 绩效考核 ===
|
|
|
|
|
+ {
|
|
|
|
|
+ "question": "绩效考核结果分几个等级?",
|
|
|
|
|
+ "section_keywords": ["绩效考核", "考核等级", "绩效"],
|
|
|
|
|
+ "reference_answer": "绩效考核分为S/A/B/C/D五个等级,其中S级占比不超过10%。",
|
|
|
|
|
+ },
|
|
|
|
|
+ # === 第七章 培训与发展 ===
|
|
|
|
|
+ {
|
|
|
|
|
+ "question": "公司每年要求的最低培训学时是多少?",
|
|
|
|
|
+ "section_keywords": ["培训", "培训与发展", "学习"],
|
|
|
|
|
+ "reference_answer": "每位员工每年需完成不低于40学时的培训,包括技术培训、管理培训和软技能培训。",
|
|
|
|
|
+ },
|
|
|
|
|
+ # === 第八章 晋升制度 ===
|
|
|
|
|
+ {
|
|
|
|
|
+ "question": "员工晋升需要满足哪些基本条件?",
|
|
|
|
|
+ "section_keywords": ["晋升", "晋升制度", "职业发展"],
|
|
|
|
|
+ "reference_answer": "晋升需满足:连续两次绩效考核B+以上,在现岗位工作满1年以上。",
|
|
|
|
|
+ },
|
|
|
|
|
+ # === 第九章 加班管理 ===
|
|
|
|
|
+ {
|
|
|
|
|
+ "question": "周末加班工资按几倍计算?",
|
|
|
|
|
+ "section_keywords": ["加班", "加班费", "加班工资"],
|
|
|
|
|
+ "reference_answer": "休息日(周末)加班,按200%(2倍)工资支付加班费。",
|
|
|
|
|
+ },
|
|
|
|
|
+ # === 第十章 出差与报销 ===
|
|
|
|
|
+ {
|
|
|
|
|
+ "question": "出差一线城市住宿标准多少钱一天?",
|
|
|
|
|
+ "section_keywords": ["出差", "报销", "住宿标准"],
|
|
|
|
|
+ "reference_answer": "一线城市(北上广深)出差住宿标准为500元/天。",
|
|
|
|
|
+ },
|
|
|
|
|
+ # === 第十一章 劳动合同 ===
|
|
|
|
|
+ {
|
|
|
|
|
+ "question": "第一次签劳动合同的期限是几年?",
|
|
|
|
|
+ "section_keywords": ["劳动合同", "合同期限", "签订"],
|
|
|
|
|
+ "reference_answer": "首次签订劳动合同期限为3年,续签第二次为5年,之后可签无固定期限合同。",
|
|
|
|
|
+ },
|
|
|
|
|
+ # === 第十二章 奖惩制度 ===
|
|
|
|
|
+ {
|
|
|
|
|
+ "question": "员工连续旷工几天会被开除?",
|
|
|
|
|
+ "section_keywords": ["旷工", "处罚", "奖惩", "开除", "严重违纪"],
|
|
|
|
|
+ "reference_answer": "连续旷工3天以上或一年内累计旷工5天以上,属于严重违纪,公司可单方解除劳动合同。",
|
|
|
|
|
+ },
|
|
|
|
|
+ # === 第十三章 保密规定 ===
|
|
|
|
|
+ {
|
|
|
|
|
+ "question": "公司商业机密包括哪些内容?",
|
|
|
|
|
+ "section_keywords": ["保密", "商业机密", "保密规定"],
|
|
|
|
|
+ "reference_answer": "商业机密包括技术资料、客户信息、财务数据、经营决策、未公开的人事信息等。",
|
|
|
|
|
+ },
|
|
|
|
|
+ # === 第十四章 离职管理 ===
|
|
|
|
|
+ {
|
|
|
|
|
+ "question": "离职需要提前多少天申请?",
|
|
|
|
|
+ "section_keywords": ["离职", "离职管理", "解除"],
|
|
|
|
|
+ "reference_answer": "正式员工离职需提前30天书面申请,试用期员工提前3天申请。",
|
|
|
|
|
+ },
|
|
|
|
|
+ # === 第十五章 员工关怀 ===
|
|
|
|
|
+ {
|
|
|
|
|
+ "question": "公司有哪些员工关怀和福利项目?",
|
|
|
|
|
+ "section_keywords": ["员工关怀", "员工福利", "关怀"],
|
|
|
|
|
+ "reference_answer": "公司提供补充商业保险、年度体检、节日礼品、团队建设、生日福利、婚育贺礼、困难补助等。",
|
|
|
|
|
+ },
|
|
|
|
|
+ # === 第十六章 附则 ===
|
|
|
|
|
+ {
|
|
|
|
|
+ "question": "员工手册的解释权归哪个部门?",
|
|
|
|
|
+ "section_keywords": ["附则", "解释权", "人力资源部"],
|
|
|
|
|
+ "reference_answer": "本手册的最终解释权归公司人力资源部所有。",
|
|
|
|
|
+ },
|
|
|
|
|
+ {
|
|
|
|
|
+ "question": "员工手册从什么时候开始生效?",
|
|
|
|
|
+ "section_keywords": ["附则", "生效", "施行"],
|
|
|
|
|
+ "reference_answer": "本手册自2025年1月1日起正式生效施行。",
|
|
|
|
|
+ },
|
|
|
|
|
+]
|
|
|
|
|
+
|
|
|
|
|
+def get_manual_test_set() -> list[dict]:
|
|
|
|
|
+ """获取手动标注的真实测试集。
|
|
|
|
|
+
|
|
|
|
|
+ 这些测试用例是基于员工手册原文内容手动编写的,用于客观评估 RAG 系统。
|
|
|
|
|
+ 与 LLM 自动生成的测试集不同,这里的问题和答案都是人工确定的,
|
|
|
|
|
+ 不存在"同一 LLM 既当选手又当裁判"的问题。
|
|
|
|
|
+ """
|
|
|
|
|
+ return _MANUAL_TEST_SET
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+# ========================================================================
|
|
|
|
|
+# 10b. Section 级检索评估(核心指标)
|
|
|
|
|
+# ========================================================================
|
|
|
|
|
+
|
|
|
|
|
+def evaluate_retrieval_section_level(test_set: list[dict]) -> dict:
|
|
|
|
|
+ """Section 级检索评估 — 客观评估检索质量。
|
|
|
|
|
+
|
|
|
|
|
+ 评估方法:
|
|
|
|
|
+ 对每个问题,使用各检索策略检索 Top-K 文档,检查检索到的文档
|
|
|
|
|
+ 是否来自员工手册的正确章节(通过章节关键词匹配)。
|
|
|
|
|
+
|
|
|
|
|
+ 为什么不用"同一个 chunk"检测?
|
|
|
|
|
+ → 那是循环验证:LLM 从 chunk A 生成问题 → 检查是否检索到 chunk A
|
|
|
|
|
+ → 毫无意义,因为 embedding 向量天然相近。
|
|
|
|
|
+
|
|
|
|
|
+ 为什么用 Section 级匹配?
|
|
|
|
|
+ → 比如问"试用期",理应检索到"第二章 入职与试用期"的内容
|
|
|
|
|
+ → 如果检索到了"第四章 假期管理",说明检索跑偏了
|
|
|
|
|
+ → 这能真正区分不同检索策略的效果
|
|
|
|
|
+
|
|
|
|
|
+ Returns:
|
|
|
|
|
+ dict: 每种策略的 {"precision@3", "precision@5", "section_recall", "details"}
|
|
|
|
|
+ """
|
|
|
|
|
+ print(f"\n{'='*60}")
|
|
|
|
|
+ print(f"[检索评估] Section 级评估 (测试 {len(test_set)} 个真实查询)")
|
|
|
|
|
+ print(f"[检索评估] 衡量检索到的文档是否来自正确章节")
|
|
|
|
|
+
|
|
|
|
|
+ strategies = ["bm25", "vector", "ensemble", "multi_query", "enhanced"]
|
|
|
|
|
+ all_results = {}
|
|
|
|
|
+
|
|
|
|
|
+ for strategy in strategies:
|
|
|
|
|
+ per_query = []
|
|
|
|
|
+
|
|
|
|
|
+ for qa in test_set:
|
|
|
|
|
+ query = qa["question"]
|
|
|
|
|
+ expected_keywords = qa["section_keywords"]
|
|
|
|
|
+
|
|
|
|
|
+ # "enhanced" 用增强检索链(重写+HyDE+RRF,是实际生产 pipeline)
|
|
|
|
|
+ if strategy == "enhanced":
|
|
|
|
|
+ docs = enhanced_retrieve(query, strategy="vector", top_k=10)
|
|
|
|
|
+ else:
|
|
|
|
|
+ docs = retrieve(query, strategy=strategy, top_k=10)
|
|
|
|
|
+
|
|
|
|
|
+ # 对每个检索到的 doc,检查其内容是否包含期望的章节关键词
|
|
|
|
|
+ top3_hits = 0
|
|
|
|
|
+ top5_hits = 0
|
|
|
|
|
+ top3_texts = []
|
|
|
|
|
+ for rank, doc in enumerate(docs):
|
|
|
|
|
+ doc_text = doc.page_content
|
|
|
|
|
+ # 检查文档中是否包含任何期望的章节关键词
|
|
|
|
|
+ matched = any(kw in doc_text for kw in expected_keywords)
|
|
|
|
|
+ if matched:
|
|
|
|
|
+ if rank < 3:
|
|
|
|
|
+ top3_hits += 1
|
|
|
|
|
+ if rank < 5:
|
|
|
|
|
+ top5_hits += 1
|
|
|
|
|
+ if rank < 3:
|
|
|
|
|
+ top3_texts.append(doc_text[:80])
|
|
|
|
|
+
|
|
|
|
|
+ per_query.append({
|
|
|
|
|
+ "query": query[:60],
|
|
|
|
|
+ "expected_sections": expected_keywords,
|
|
|
|
|
+ "top3_hits": top3_hits,
|
|
|
|
|
+ "top5_hits": top5_hits,
|
|
|
|
|
+ "top3_preview": top3_texts,
|
|
|
|
|
+ })
|
|
|
|
|
+
|
|
|
|
|
+ n = len(test_set)
|
|
|
|
|
+ precision_3 = sum(q["top3_hits"] for q in per_query) / (n * 3) # 理想情况是 top3 中每篇都命中
|
|
|
|
|
+ precision_5 = sum(q["top5_hits"] for q in per_query) / (n * 5)
|
|
|
|
|
+ # Section Recall: 至少有一篇文档命中正确章节的查询比例
|
|
|
|
|
+ section_recall = sum(1 for q in per_query if q["top3_hits"] > 0) / n
|
|
|
|
|
+
|
|
|
|
|
+ all_results[strategy] = {
|
|
|
|
|
+ "precision_at_3": precision_3,
|
|
|
|
|
+ "precision_at_5": precision_5,
|
|
|
|
|
+ "section_recall": section_recall,
|
|
|
|
|
+ "details": per_query,
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ print(f"[检索评估] {strategy:<15s} | P@3: {precision_3:.1%} | "
|
|
|
|
|
+ f"P@5: {precision_5:.1%} | Section Recall: {section_recall:.1%}")
|
|
|
|
|
+
|
|
|
|
|
+ return all_results
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+# ========================================================================
|
|
|
|
|
+# 10c. 检索交叉验证
|
|
|
|
|
+# ========================================================================
|
|
|
|
|
+
|
|
|
|
|
+def evaluate_retrieval_cross_verify(test_set: list[dict]) -> dict:
|
|
|
|
|
+ """检索交叉验证: LLM 逐条判断检索到的文档是否能回答问题。
|
|
|
|
|
+
|
|
|
|
|
+ 与 Section 级匹配互补:
|
|
|
|
|
+ - Section 匹配是"粗粒度"的(文档来自正确章节?)
|
|
|
|
|
+ - 交叉验证是"细粒度"的(文档内容真的能回答这个问题?)
|
|
|
|
|
+
|
|
|
|
|
+ 方法:
|
|
|
|
|
+ 对每个查询,取检索到的 Top-3 文档,让 LLM 逐条判断:
|
|
|
|
|
+ "这篇文档是否包含回答该问题所需的信息?"
|
|
|
|
|
+ → 是(1) / 部分(0.5) / 否(0)
|
|
|
|
|
+
|
|
|
|
|
+ 优势: 不需要人工标注参考答案,LLM 只做"是否相关"的二元判断,
|
|
|
|
|
+ 比对比参考答案更客观。
|
|
|
|
|
+ """
|
|
|
|
|
+ print(f"\n{'='*60}")
|
|
|
|
|
+ print(f"[交叉验证] LLM 逐条判断检索文档是否相关")
|
|
|
|
|
+
|
|
|
|
|
+ strategies = ["bm25", "vector", "ensemble", "multi_query"]
|
|
|
|
|
+ all_results = {}
|
|
|
|
|
+ llm = get_llm()
|
|
|
|
|
+
|
|
|
|
|
+ for strategy in strategies:
|
|
|
|
|
+ doc_scores = []
|
|
|
|
|
+
|
|
|
|
|
+ for qi, qa in enumerate(test_set):
|
|
|
|
|
+ query = qa["question"]
|
|
|
|
|
+ docs = retrieve(query, strategy=strategy, top_k=3)
|
|
|
|
|
+
|
|
|
|
|
+ for rank, doc in enumerate(docs):
|
|
|
|
|
+ verify_prompt = f"""判断以下文档片段是否包含回答该问题所需的信息。
|
|
|
|
|
+
|
|
|
|
|
+问题:{query}
|
|
|
|
|
+文档:{doc.page_content[:400]}
|
|
|
|
|
+
|
|
|
|
|
+只回答数字: 1=包含关键信息 / 0.5=部分相关 / 0=无关"""
|
|
|
|
|
+
|
|
|
|
|
+ try:
|
|
|
|
|
+ raw = llm.invoke(verify_prompt)
|
|
|
|
|
+ ans = raw.content if hasattr(raw, "content") else str(raw)
|
|
|
|
|
+ score = float(ans.strip()[0]) if ans.strip() else 0
|
|
|
|
|
+ score = max(0, min(1, score))
|
|
|
|
|
+ except Exception:
|
|
|
|
|
+ score = 0
|
|
|
|
|
+
|
|
|
|
|
+ doc_scores.append({
|
|
|
|
|
+ "query": query[:50], "strategy": strategy,
|
|
|
|
|
+ "rank": rank + 1, "relevant": score,
|
|
|
|
|
+ })
|
|
|
|
|
+
|
|
|
|
|
+ n = len(test_set) * 3 # 每个查询 3 篇文档
|
|
|
|
|
+ avg_relevance = sum(d["relevant"] for d in doc_scores) / max(1, n)
|
|
|
|
|
+ perfect_ratio = sum(1 for d in doc_scores if d["relevant"] == 1) / max(1, n)
|
|
|
|
|
+
|
|
|
|
|
+ all_results[strategy] = {
|
|
|
|
|
+ "avg_relevance": avg_relevance,
|
|
|
|
|
+ "perfect_ratio": perfect_ratio,
|
|
|
|
|
+ "details": doc_scores,
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ print(f"[交叉验证] {strategy:<15s} | 平均相关度: {avg_relevance:.2%} | "
|
|
|
|
|
+ f"完美相关比: {perfect_ratio:.1%}")
|
|
|
|
|
+
|
|
|
|
|
+ return all_results
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+# ========================================================================
|
|
|
|
|
+# 10d. 生成质量评估 + Bad Case 分析
|
|
|
|
|
+# ========================================================================
|
|
|
|
|
+
|
|
|
|
|
+def evaluate_generation(test_set: list[dict], sample_size: int = 10) -> dict:
|
|
|
|
|
+ """生成效果评估: 准确率、完整度、相关度、幻觉率。
|
|
|
|
|
+
|
|
|
|
|
+ 对每个测试问题:
|
|
|
|
|
+ 1. RAG 生成答案
|
|
|
|
|
+ 2. LLM judge 对比"参考答案"和"生成答案"
|
|
|
|
|
+ 3. 从 4 个维度打分
|
|
|
|
|
+
|
|
|
|
|
+ 注意: 这个评估的可靠性取决于 LLM judge 的质量,建议作为辅助指标。
|
|
|
|
|
+ """
|
|
|
|
|
+ import random
|
|
|
|
|
+
|
|
|
|
|
+ print(f"\n{'='*60}")
|
|
|
|
|
+ print(f"[生成评估] LLM Judge 打分(辅助指标)")
|
|
|
|
|
+
|
|
|
|
|
+ if len(test_set) > sample_size:
|
|
|
|
|
+ samples = random.sample(test_set, sample_size)
|
|
|
|
|
+ else:
|
|
|
|
|
+ samples = test_set
|
|
|
|
|
+
|
|
|
|
|
+ llm = get_llm()
|
|
|
|
|
+ judge_results = []
|
|
|
|
|
+ bad_cases = []
|
|
|
|
|
+
|
|
|
|
|
+ for i, qa in enumerate(samples):
|
|
|
|
|
+ query = qa["question"]
|
|
|
|
|
+ reference = qa["reference_answer"]
|
|
|
|
|
+
|
|
|
|
|
+ # RAG 生成答案
|
|
|
|
|
+ rag_result = rag_query(query, use_enhance=True, use_rerank=True)
|
|
|
|
|
+ generated = rag_result["answer"]
|
|
|
|
|
+
|
|
|
|
|
+ # 同时获取检索结果,用于 Bad Case 根因分析
|
|
|
|
|
+ retrieved_docs = retrieve(query, strategy="ensemble", top_k=5)
|
|
|
|
|
+
|
|
|
|
|
+ # LLM judge
|
|
|
|
|
+ judge_prompt = f"""你是一个严格但公正的评估助手。请对比参考答案和RAG系统生成的答案。
|
|
|
|
|
+
|
|
|
|
|
+参考答案(来自员工手册原文):{reference}
|
|
|
|
|
+RAG生成答案:{generated}
|
|
|
|
|
+
|
|
|
|
|
+请从以下维度严格打分(1-5分,5分最好):
|
|
|
|
|
+- accuracy(事实准确性):生成答案的事实是否与参考答案一致?(有事实错误扣分)
|
|
|
|
|
+- completeness(完整性):是否覆盖了参考答案的关键信息?(遗漏重要信息扣分)
|
|
|
|
|
+- relevance(相关性):是否直接回答了问题?(答非所问扣分)
|
|
|
|
|
+- hallucination(幻觉程度):是否编造了参考答案中没有的内容?(5=无编造,1=严重编造)
|
|
|
|
|
+
|
|
|
|
|
+输出严格的JSON:
|
|
|
|
|
+{{"accuracy": 整数1-5, "completeness": 整数1-5, "relevance": 整数1-5, "hallucination": 整数1-5}}"""
|
|
|
|
|
+
|
|
|
|
|
+ try:
|
|
|
|
|
+ raw_resp = llm.invoke(judge_prompt)
|
|
|
|
|
+ raw_text = raw_resp.content if hasattr(raw_resp, "content") else str(raw_resp)
|
|
|
|
|
+ json_match = re.search(r'\{[\s\S]*\}', raw_text)
|
|
|
|
|
+ if json_match:
|
|
|
|
|
+ scores = json.loads(json_match.group())
|
|
|
|
|
+ scores["query"] = query
|
|
|
|
|
+ scores["reference"] = reference
|
|
|
|
|
+ scores["generated"] = generated[:300]
|
|
|
|
|
+ judge_results.append(scores)
|
|
|
|
|
+
|
|
|
|
|
+ # 识别 Bad Case
|
|
|
|
|
+ is_bad = (scores.get("accuracy", 5) <= 2 or
|
|
|
|
|
+ scores.get("hallucination", 5) <= 2)
|
|
|
|
|
+ if is_bad:
|
|
|
|
|
+ # 根因分析
|
|
|
|
|
+ retrieved_texts = [d.page_content[:200] for d in retrieved_docs[:3]]
|
|
|
|
|
+ # 检查检索结果中是否有与 reference 相关的内容
|
|
|
|
|
+ retrieval_ok = any(
|
|
|
|
|
+ any(kw in " ".join(retrieved_texts) for kw in qa.get("section_keywords", []))
|
|
|
|
|
+ )
|
|
|
|
|
+ root_cause = "生成幻觉" if retrieval_ok else "检索失败"
|
|
|
|
|
+ bad_cases.append({
|
|
|
|
|
+ "query": query,
|
|
|
|
|
+ "reference": reference,
|
|
|
|
|
+ "generated": generated[:200],
|
|
|
|
|
+ "root_cause": root_cause,
|
|
|
|
|
+ "retrieved_preview": retrieved_texts[:2],
|
|
|
|
|
+ "scores": scores,
|
|
|
|
|
+ })
|
|
|
|
|
+
|
|
|
|
|
+ print(f"[生成评估] [{i+1}/{len(samples)}] "
|
|
|
|
|
+ f"acc={scores.get('accuracy','?')} "
|
|
|
|
|
+ f"hall={scores.get('hallucination','?')} "
|
|
|
|
|
+ f"{'⚠️ BAD' if is_bad else '✅'}")
|
|
|
|
|
+
|
|
|
|
|
+ except Exception as e:
|
|
|
|
|
+ print(f"[生成评估] 第 {i+1} 个评判失败: {e}")
|
|
|
|
|
+
|
|
|
|
|
+ # 汇总指标
|
|
|
|
|
+ if judge_results:
|
|
|
|
|
+ avg_acc = sum(r.get("accuracy", 0) for r in judge_results) / len(judge_results)
|
|
|
|
|
+ avg_comp = sum(r.get("completeness", 0) for r in judge_results) / len(judge_results)
|
|
|
|
|
+ avg_rel = sum(r.get("relevance", 0) for r in judge_results) / len(judge_results)
|
|
|
|
|
+ avg_hall = sum(r.get("hallucination", 0) for r in judge_results) / len(judge_results)
|
|
|
|
|
+ hall_rate = sum(1 for r in judge_results if r.get("hallucination", 5) <= 2) / len(judge_results)
|
|
|
|
|
+ bad_rate = len(bad_cases) / len(judge_results)
|
|
|
|
|
+ else:
|
|
|
|
|
+ avg_acc = avg_comp = avg_rel = avg_hall = hall_rate = bad_rate = 0
|
|
|
|
|
+
|
|
|
|
|
+ print(f"\n[生成评估] 汇总:")
|
|
|
|
|
+ print(f" 准确率: {avg_acc:.2f}/5 | 完整度: {avg_comp:.2f}/5")
|
|
|
|
|
+ print(f" 相关度: {avg_rel:.2f}/5 | 抗幻觉: {avg_hall:.2f}/5")
|
|
|
|
|
+ print(f" 幻觉率: {hall_rate:.1%} | Bad率: {bad_rate:.1%}")
|
|
|
|
|
+
|
|
|
|
|
+ return {
|
|
|
|
|
+ "avg_accuracy": avg_acc,
|
|
|
|
|
+ "avg_completeness": avg_comp,
|
|
|
|
|
+ "avg_relevance": avg_rel,
|
|
|
|
|
+ "avg_anti_hallucination": avg_hall,
|
|
|
|
|
+ "hallucination_rate": hall_rate,
|
|
|
|
|
+ "bad_case_rate": bad_rate,
|
|
|
|
|
+ "bad_cases": bad_cases,
|
|
|
|
|
+ "details": judge_results,
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def bad_case_analysis(gen_results: dict) -> list[dict]:
|
|
|
|
|
+ """Bad Case 深度分析与优化建议。
|
|
|
|
|
+
|
|
|
|
|
+ 对每个 Bad Case 按根因分类,给出针对性优化方向。
|
|
|
|
|
+ """
|
|
|
|
|
+ bad_cases = gen_results.get("bad_cases", [])
|
|
|
|
|
+
|
|
|
|
|
+ print(f"\n{'='*60}")
|
|
|
|
|
+ print(f"[Bad Case 分析] 根因定位")
|
|
|
|
|
+
|
|
|
|
|
+ if not bad_cases:
|
|
|
|
|
+ print("[Bad Case 分析] ✅ 所有案例表现良好!")
|
|
|
|
|
+ return []
|
|
|
|
|
+
|
|
|
|
|
+ # 按根因分类
|
|
|
|
|
+ categories = defaultdict(list)
|
|
|
|
|
+ for case in bad_cases:
|
|
|
|
|
+ categories[case["root_cause"]].append(case)
|
|
|
|
|
+
|
|
|
|
|
+ print(f"[Bad Case 分析] 共 {len(bad_cases)} 个 Bad Case:\n")
|
|
|
|
|
+
|
|
|
|
|
+ for root_cause, cases in sorted(categories.items()):
|
|
|
|
|
+ print(f" 📌 [{root_cause}] ({len(cases)} 例)")
|
|
|
|
|
+ for case in cases[:1]:
|
|
|
|
|
+ print(f" ❓ {case['query'][:60]}")
|
|
|
|
|
+ print(f" 🤖 {case['generated'][:80]}...")
|
|
|
|
|
+ print()
|
|
|
|
|
+
|
|
|
|
|
+ # 优化建议
|
|
|
|
|
+ print(f"[Bad Case 分析] 针对性优化建议:")
|
|
|
|
|
+ if "检索失败" in categories:
|
|
|
|
|
+ print(f" 🔧 检索失败 ({len(categories['检索失败'])}例):")
|
|
|
|
|
+ print(f" → 降低 SemanticChunker 的 percentile 值(更细粒度切块)")
|
|
|
|
|
+ print(f" → 增加 RETRIEVAL_K 召回量")
|
|
|
|
|
+ print(f" → 考虑换用更大模型或调整 chunk 粒度")
|
|
|
|
|
+ print(f" → 检查查询重写是否改变了语义")
|
|
|
|
|
+ if "生成幻觉" in categories:
|
|
|
|
|
+ print(f" 🔧 生成幻觉 ({len(categories['生成幻觉'])}例):")
|
|
|
|
|
+ print(f" → 在 prompt 中加强对「不知道就说不知道」的约束")
|
|
|
|
|
+ print(f" → 降低 LLM temperature(当前默认值)")
|
|
|
|
|
+ print(f" → 要求逐句标注引用来源")
|
|
|
|
|
+ print(f" → 考虑接入 Reranker 精排后再生成")
|
|
|
|
|
+
|
|
|
|
|
+ return bad_cases
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+# ========================================================================
|
|
|
|
|
+# 10e. 完整评估入口
|
|
|
|
|
+# ========================================================================
|
|
|
|
|
+
|
|
|
|
|
+def run_full_evaluation() -> dict:
|
|
|
|
|
+ """运行完整评估流程。
|
|
|
|
|
+
|
|
|
|
|
+ 三层评估:
|
|
|
|
|
+ 第一层: Section 级检索评估 → 客观指标(不依赖 LLM judge)
|
|
|
|
|
+ 第二层: 检索交叉验证 → LLM 逐篇判断相关性
|
|
|
|
|
+ 第三层: 生成质量评估 → LLM judge 打分 + Bad Case 根因分析
|
|
|
|
|
+
|
|
|
|
|
+ Returns:
|
|
|
|
|
+ dict: 完整评估报告
|
|
|
|
|
+ """
|
|
|
|
|
+ print(f"\n{'#'*60}")
|
|
|
|
|
+ print(f"# RAG 效果评估(三层体系)")
|
|
|
|
|
+ print(f"{'#'*60}")
|
|
|
|
|
+
|
|
|
|
|
+ # 使用手动标注的真实测试集(覆盖全部 16 章)
|
|
|
|
|
+ test_set = get_manual_test_set()
|
|
|
|
|
+ print(f"[评估] 手动标注测试集: {len(test_set)} 条,覆盖全部 16 章")
|
|
|
|
|
+
|
|
|
|
|
+ # ---- 第一层: Section 级检索评估 ----
|
|
|
|
|
+ section_results = evaluate_retrieval_section_level(test_set)
|
|
|
|
|
+
|
|
|
|
|
+ # ---- 第二层: 检索交叉验证 ----
|
|
|
|
|
+ cross_verify_results = evaluate_retrieval_cross_verify(test_set)
|
|
|
|
|
+
|
|
|
|
|
+ # ---- 第三层: 生成质量评估 + Bad Case ----
|
|
|
|
|
+ gen_results = evaluate_generation(test_set, sample_size=min(12, len(test_set)))
|
|
|
|
|
+ bad_cases = bad_case_analysis(gen_results)
|
|
|
|
|
+
|
|
|
|
|
+ # ---- 综合报告 ----
|
|
|
|
|
+ print(f"\n{'='*60}")
|
|
|
|
|
+ print(f"📊 综合评估报告")
|
|
|
|
|
+ print(f"{'='*60}")
|
|
|
|
|
+ print(f"测试集: {len(test_set)} 条手动标注的真实员工问题\n")
|
|
|
|
|
+
|
|
|
|
|
+ # 检索排名
|
|
|
|
|
+ print(f"🏆 检索策略排名 (按 Section Recall):")
|
|
|
|
|
+ sorted_sr = sorted(section_results.items(),
|
|
|
|
|
+ key=lambda x: x[1]["section_recall"], reverse=True)
|
|
|
|
|
+ for rank, (name, m) in enumerate(sorted_sr):
|
|
|
|
|
+ medal = ["🥇", "🥈", "🥉", "4️⃣", "5️⃣"][rank] if rank < 5 else " "
|
|
|
|
|
+ print(f" {medal} {name:<15s} P@3: {m['precision_at_3']:.1%} "
|
|
|
|
|
+ f"P@5: {m['precision_at_5']:.1%} "
|
|
|
|
|
+ f"SectionRecall: {m['section_recall']:.1%}")
|
|
|
|
|
+
|
|
|
|
|
+ print(f"\n🔍 检索交叉验证 (LLM 逐篇判断相关性):")
|
|
|
|
|
+ sorted_cv = sorted(cross_verify_results.items(),
|
|
|
|
|
+ key=lambda x: x[1]["avg_relevance"], reverse=True)
|
|
|
|
|
+ for rank, (name, m) in enumerate(sorted_cv):
|
|
|
|
|
+ medal = ["🥇", "🥈", "🥉", "4️⃣", "5️⃣"][rank] if rank < 5 else " "
|
|
|
|
|
+ print(f" {medal} {name:<15s} 平均相关度: {m['avg_relevance']:.2%} "
|
|
|
|
|
+ f"完美相关比: {m['perfect_ratio']:.1%}")
|
|
|
|
|
+
|
|
|
|
|
+ print(f"\n📈 生成质量 (LLM Judge 辅助指标):")
|
|
|
|
|
+ print(f" 准确率: {gen_results.get('avg_accuracy', 0):.2f}/5")
|
|
|
|
|
+ print(f" 完整度: {gen_results.get('avg_completeness', 0):.2f}/5")
|
|
|
|
|
+ print(f" 幻觉率: {gen_results.get('hallucination_rate', 0):.1%}")
|
|
|
|
|
+ print(f" Bad率: {gen_results.get('bad_case_rate', 0):.1%}")
|
|
|
|
|
+ print(f" Bad Case: {len(bad_cases)} 个")
|
|
|
|
|
+
|
|
|
|
|
+ # 确定最佳策略
|
|
|
|
|
+ best_strategy = sorted_sr[0][0] if sorted_sr else "ensemble"
|
|
|
|
|
+ best_cv = sorted_cv[0][0] if sorted_cv else "ensemble"
|
|
|
|
|
+ print(f"\n💡 推荐策略: 检索={best_strategy}, 交叉验证最佳={best_cv}")
|
|
|
|
|
+ if best_strategy != best_cv:
|
|
|
|
|
+ print(f" ⚠️ Section评估与交叉验证结论不一致,建议综合参考")
|
|
|
|
|
+
|
|
|
|
|
+ return {
|
|
|
|
|
+ "test_set_size": len(test_set),
|
|
|
|
|
+ "section_retrieval": section_results,
|
|
|
|
|
+ "cross_verify": cross_verify_results,
|
|
|
|
|
+ "generation": gen_results,
|
|
|
|
|
+ "bad_cases": bad_cases,
|
|
|
|
|
+ "recommended_strategy": best_strategy,
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+# ============================================================================
|
|
|
|
|
+# Section 11: 主流程
|
|
|
|
|
+# ============================================================================
|
|
|
|
|
+
|
|
|
|
|
+def _print_banner():
|
|
|
|
|
+ print("""
|
|
|
|
|
+╔══════════════════════════════════════════════════════════╗
|
|
|
|
|
+║ RAG 全链路优化系统 — 员工手册智能问答 ║
|
|
|
|
|
+║ ║
|
|
|
|
|
+║ 1.员工手册.md 2.Agent工具 3.BGE-Large-Zh(1024d) ║
|
|
|
|
|
+║ 4.语义切块 5.多路召回 6.查询增强(HyDE) ║
|
|
|
|
|
+║ 7.Reranker精排 8.三层评估体系 ║
|
|
|
|
|
+╚══════════════════════════════════════════════════════════╝
|
|
|
|
|
+""")
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+def main():
|
|
|
|
|
+ """主流程: 初始化 → 示例验证 → 交互测试。
|
|
|
|
|
+
|
|
|
|
|
+ 优化:默认跳过完整评估(200+ 次 LLM 调用,耗时 10+ 分钟),
|
|
|
|
|
+ 改为运行 1 个示例查询快速验证 RAG 链路可用,再进入交互模式。
|
|
|
|
|
+ 完整评估通过命令行参数 `--eval` 或交互命令 `eval` 触发。
|
|
|
|
|
+ """
|
|
|
|
|
+ global _VS, _CHUNKS
|
|
|
|
|
+
|
|
|
|
|
+ # 命令行参数: --eval 启动时直接跑完整评估;--demo 只跑示例不进入交互
|
|
|
|
|
+ run_eval_at_startup = "--eval" in sys.argv
|
|
|
|
|
+ demo_only = "--demo" in sys.argv
|
|
|
|
|
+
|
|
|
|
|
+ _print_banner()
|
|
|
|
|
+
|
|
|
|
|
+ # ---- Step 1: 核心初始化(仅加载必要组件)----
|
|
|
|
|
+ print("[初始化] 加载模型和文档...\n")
|
|
|
|
|
+ init_embedding_model() # embedding 模型(必须)
|
|
|
|
|
+ _CHUNKS = load_and_chunk() # 文档切块(必须)
|
|
|
|
|
+ _VS = get_vectorstore(_CHUNKS) # 向量存储(必须)
|
|
|
|
|
+ create_ensemble_retriever(_CHUNKS, _VS) # 默认策略检索器(必须)
|
|
|
|
|
+
|
|
|
|
|
+ # 以下组件改为按需初始化,避免启动时加载不必要的模型
|
|
|
|
|
+ # - create_mq_retriever(): 仅 multi_query 策略用到,首次调用时再初始化
|
|
|
|
|
+ # - init_reranker(): CrossEncoder 1.1GB,加载慢,rerank() 内部已按需加载
|
|
|
|
|
+
|
|
|
|
|
+ # ---- Step 2: 快速示例验证(1 个查询,验证 RAG 链路可用)----
|
|
|
|
|
+ print(f"\n{'='*60}")
|
|
|
|
|
+ print(f"[示例] 快速验证 RAG 链路(不启用查询增强/精排,约 3 秒)")
|
|
|
|
|
+ print(f"{'='*60}")
|
|
|
|
|
+ demo_q = "试用期是几个月?"
|
|
|
|
|
+ print(f"问题: {demo_q}")
|
|
|
|
|
+ result = rag_query(demo_q, strategy="ensemble",
|
|
|
|
|
+ use_enhance=False, use_rerank=False)
|
|
|
|
|
+ print(f"答案: {result['answer'][:200]}")
|
|
|
|
|
+ print(f"来源: {len(result['sources'])} 个文档 "
|
|
|
|
|
+ f"(top-1: {result['sources'][0].metadata.get('section', '?')})")
|
|
|
|
|
+
|
|
|
|
|
+ # ---- Step 3: 可选完整评估(默认跳过)----
|
|
|
|
|
+ if run_eval_at_startup:
|
|
|
|
|
+ eval_report = run_full_evaluation()
|
|
|
|
|
+ else:
|
|
|
|
|
+ print(f"\n[提示] 完整评估已跳过(耗时较长)。")
|
|
|
|
|
+ print(f" 如需运行评估:交互模式输入 'eval',"
|
|
|
|
|
+ f"或重启时加 --eval 参数")
|
|
|
|
|
+
|
|
|
|
|
+ # ---- Step 4: 交互模式 ----
|
|
|
|
|
+ if demo_only:
|
|
|
|
|
+ print(f"\n[--demo] 示例完成,不进入交互模式")
|
|
|
|
|
+ return
|
|
|
|
|
+
|
|
|
|
|
+ print(f"\n{'='*60}")
|
|
|
|
|
+ print(f"💬 交互模式")
|
|
|
|
|
+ print(f"{'='*60}")
|
|
|
|
|
+ print(f" 输入问题开始查询")
|
|
|
|
|
+ print(f" 命令: 'strategy <name>' 切换检索策略")
|
|
|
|
|
+ print(f" 策略: bm25 | vector | ensemble | multi_query")
|
|
|
|
|
+ print(f" 命令: 'agent <query>' 使用 Agent 模式")
|
|
|
|
|
+ print(f" 命令: 'enhance on/off' 切换查询增强(重写+HyDE)")
|
|
|
|
|
+ print(f" 命令: 'rerank on/off' 切换精排")
|
|
|
|
|
+ print(f" 命令: 'eval' 运行完整评估(耗时)")
|
|
|
|
|
+ print(f" 命令: 'exit' 退出")
|
|
|
|
|
+ print(f"{'='*60}\n")
|
|
|
|
|
+
|
|
|
|
|
+ current_strategy = "ensemble"
|
|
|
|
|
+ use_enhance = False # 默认关闭(每次查询省 2 次 LLM 调用)
|
|
|
|
|
+ use_rerank = False # 默认关闭(首次调用时加载 Reranker 模型)
|
|
|
|
|
+
|
|
|
|
|
+ while True:
|
|
|
|
|
+ try:
|
|
|
|
|
+ user_input = input("🔍 > ").strip()
|
|
|
|
|
+ except (EOFError, KeyboardInterrupt):
|
|
|
|
|
+ print("\n再见!")
|
|
|
|
|
+ break
|
|
|
|
|
+
|
|
|
|
|
+ if not user_input:
|
|
|
|
|
+ continue
|
|
|
|
|
+
|
|
|
|
|
+ if user_input.lower() in ("exit", "quit", "q"):
|
|
|
|
|
+ print("再见!")
|
|
|
|
|
+ break
|
|
|
|
|
+
|
|
|
|
|
+ # 命令处理
|
|
|
|
|
+ if user_input.lower().startswith("strategy "):
|
|
|
|
|
+ new_strategy = user_input.split(" ", 1)[1].strip()
|
|
|
|
|
+ if new_strategy in ("bm25", "vector", "ensemble", "multi_query"):
|
|
|
|
|
+ current_strategy = new_strategy
|
|
|
|
|
+ print(f"✅ 已切换到: {current_strategy}")
|
|
|
|
|
+ else:
|
|
|
|
|
+ print(f"❌ 未知策略: {new_strategy}")
|
|
|
|
|
+ continue
|
|
|
|
|
+
|
|
|
|
|
+ if user_input.lower().startswith("agent "):
|
|
|
|
|
+ query = user_input.split(" ", 1)[1].strip()
|
|
|
|
|
+ print(f"🤖 Agent 模式(自动判断是否需要检索)...\n")
|
|
|
|
|
+ answer = agent_query(query)
|
|
|
|
|
+ print(f"\n📝 {answer}\n")
|
|
|
|
|
+ continue
|
|
|
|
|
+
|
|
|
|
|
+ if user_input.lower().startswith("enhance "):
|
|
|
|
|
+ arg = user_input.split(" ", 1)[1].strip().lower()
|
|
|
|
|
+ use_enhance = arg in ("on", "true", "1", "yes")
|
|
|
|
|
+ print(f"✅ 查询增强: {'开启' if use_enhance else '关闭'}")
|
|
|
|
|
+ continue
|
|
|
|
|
+
|
|
|
|
|
+ if user_input.lower().startswith("rerank "):
|
|
|
|
|
+ arg = user_input.split(" ", 1)[1].strip().lower()
|
|
|
|
|
+ use_rerank = arg in ("on", "true", "1", "yes")
|
|
|
|
|
+ print(f"✅ 精排: {'开启' if use_rerank else '关闭'}")
|
|
|
|
|
+ continue
|
|
|
|
|
+
|
|
|
|
|
+ if user_input.lower() == "eval":
|
|
|
|
|
+ eval_report = run_full_evaluation()
|
|
|
|
|
+ continue
|
|
|
|
|
+
|
|
|
|
|
+ # 正常查询
|
|
|
|
|
+ enhance_tag = "查询增强" if use_enhance else "无增强"
|
|
|
|
|
+ rerank_tag = "精排" if use_rerank else "无精排"
|
|
|
|
|
+ print(f"🔎 检索策略: {current_strategy} | {enhance_tag} | {rerank_tag}")
|
|
|
|
|
+ result = rag_query(user_input, strategy=current_strategy,
|
|
|
|
|
+ use_enhance=use_enhance, use_rerank=use_rerank)
|
|
|
|
|
+ print(f"\n📝 {result['answer']}\n")
|
|
|
|
|
+ print(f"📖 参考来源 ({len(result['sources'])} 个文档):")
|
|
|
|
|
+ for i, doc in enumerate(result["sources"]):
|
|
|
|
|
+ print(f" [{i+1}] {doc.page_content[:100]}...")
|
|
|
|
|
+ print()
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+if __name__ == "__main__":
|
|
|
|
|
+ main()
|