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