检索前处理:查询构建、查询翻译、查询路由
自己搞的简单代码,便于理解
1 import logging 2 from typing import List, Tuple, Literal 3 from langchain_chroma import Chroma 4 from langchain_deepseek import ChatDeepSeek 5 from langchain_huggingface import HuggingFaceEmbeddings 6 from langchain_text_splitters import RecursiveCharacterTextSplitter 7 from langchain.retrievers.multi_query import MultiQueryRetriever 8 from langchain_core.prompts import PromptTemplate, ChatPromptTemplate 9 from langchain_core.output_parsers import LineListOutputParser, StrOutputParser 10 from langchain_core.pydantic_v1 import BaseModel, Field 11 12 # ===================== 基础日志与全局配置 ===================== 13 logging.basicConfig(level=logging.INFO, format="%(message)s") 14 logger = logging.getLogger("retrieval_pipeline") 15 logging.getLogger("langchain.retrievers.multi_query").setLevel(logging.WARNING) 16 17 # 全局默认数据源(关闭路由时使用,与原有流水线完全兼容) 18 DEFAULT_DATASOURCE = "game_setting" 19 20 # ===================== 1. 基础组件初始化 ===================== 21 llm = ChatDeepSeek(model="deepseek-chat", temperature=0) 22 line_parser = LineListOutputParser() 23 str_parser = StrOutputParser() 24 embed_model = HuggingFaceEmbeddings(model_name="BAAI/bge-small-zh") 25 26 # ===================== 2. 多数据源知识库构建 ===================== 27 def build_vectorstore(docs_content: list, collection_name: str): 28 """构建独立向量库集合""" 29 text_splitter = RecursiveCharacterTextSplitter(chunk_size=300, chunk_overlap=0) 30 splits = text_splitter.create_documents(docs_content) 31 return Chroma.from_documents( 32 documents=splits, 33 embedding=embed_model, 34 collection_name=collection_name 35 ) 36 37 # 2.1 各领域独立知识库内容(可替换为你的真实文档) 38 setting_docs = [ 39 "剑客职业10级核心技能为基础剑式和格挡,普陀山关卡守关BOSS为黑风妖王,主打高物理单体伤害。", 40 "法师职业主打远程法术输出,核心技能为火球术和冰墙,新手推荐先点满火球术提升清怪效率。", 41 "游戏内装备分为白、绿、蓝、紫、橙五个品质,品质越高属性加成越高,可通过副本掉落获取。" 42 ] 43 44 story_docs = [ 45 "主角猢狲生于花果山,自名孙悟空,拜师菩提祖师习得七十二变与筋斗云。", 46 "孙悟空大闹天宫后被如来佛祖压于五行山下,五百年后被唐僧救出,踏上西天取经之路。", 47 "普陀山一难中,唐僧被黑风妖王掳走,孙悟空前往黑风洞讨要,与黑熊精打斗数回合不分胜负。" 48 ] 49 50 service_docs = [ 51 "游戏充值未到账请提供订单号和账号ID,联系客服24小时内处理补发。", 52 "账号被盗请提供注册手机号和实名认证信息,提交申诉后3个工作日内回复。", 53 "游戏闪退请更新显卡驱动,降低画质设置,若仍有问题请提交设备型号和闪退日志。" 54 ] 55 56 # 2.2 构建向量库与对应检索器映射表 57 # 向量库映射(用于HyDE检索) 58 vectorstore_map = { 59 "game_setting": build_vectorstore(setting_docs, "game_setting"), 60 "game_story": build_vectorstore(story_docs, "game_story"), 61 "customer_service": build_vectorstore(service_docs, "customer_service") 62 } 63 64 # MultiQuery检索器映射(用于多查询召回) 65 multiquery_retriever_map = { 66 key: MultiQueryRetriever.from_llm( 67 retriever=vs.as_retriever(search_kwargs={"k": 2}), 68 llm=llm 69 ) 70 for key, vs in vectorstore_map.items() 71 } 72 73 # ===================== 3. 模块零:逻辑路由(入口分发) ===================== 74 # 3.1 路由结构化输出模型 75 class RouteQuery(BaseModel): 76 """将用户问题路由到最适合回答的数据源""" 77 datasource: Literal["game_setting", "game_story", "customer_service"] = Field( 78 description="根据用户问题,选择最匹配的一个数据源,只能从给定选项中选择" 79 ) 80 81 # 3.2 路由提示词与路由链 82 route_system_prompt = """你是专业的查询路由专家,请根据用户问题的内容,路由到最合适的数据源: 83 - game_setting:游戏玩法、职业技能、关卡机制、装备系统、数值设定等游戏设定类问题 84 - game_story:游戏剧情、人物关系、故事背景、角色经历等剧情类问题 85 - customer_service:充值、账号异常、bug反馈、闪退、登录失败等客服与技术问题 86 87 只需选择一个最匹配的数据源即可。""" 88 89 route_prompt = ChatPromptTemplate.from_messages([ 90 ("system", route_system_prompt), 91 ("human", "{question}"), 92 ]) 93 94 # 绑定结构化输出,自动解析为RouteQuery对象 95 structured_llm_router = llm.with_structured_output(RouteQuery) 96 router_chain = route_prompt | structured_llm_router 97 98 def run_routing(question: str) -> str: 99 """执行路由分类,返回数据源key,失败时返回默认数据源兜底""" 100 try: 101 route_result = router_chain.invoke({"question": question}) 102 target = route_result.datasource 103 logger.info(f"[路由阶段] 分类结果:→ {target} 知识库") 104 return target 105 except Exception as e: 106 logger.warning(f"[路由阶段] 路由失败,使用默认数据源兜底:{str(e)}") 107 return DEFAULT_DATASOURCE 108 109 # ===================== 4. 模块一:查询澄清 ===================== 110 CLARIFY_GENERATE_PROMPT = PromptTemplate( 111 input_variables=["question", "max_clarify_num"], 112 template="""你是专业的查询澄清助手。请判断用户的问题是否存在语义模糊、信息缺失、范围不明确、存在歧义的情况。 113 规则: 114 1. 如果问题足够明确,不需要补充信息,请直接输出空行,不要输出任何其他内容。 115 2. 如果问题需要澄清,请生成 {max_clarify_num} 个针对性的澄清问题,用于补全信息、缩小范围、消除歧义。 116 3. 每个澄清问题只询问一个信息点,表述简洁中立。 117 4. 每个问题单独占一行,不要编号,不要多余解释。 118 119 用户原始问题:{question}""" 120 ) 121 122 QUERY_INTEGRATE_PROMPT = PromptTemplate( 123 input_variables=["original_question", "clarify_history"], 124 template="""请根据用户的原始问题和澄清问答记录,整合出一个信息完整、表述精准的最终查询语句。 125 要求: 126 1. 保留原始问题的核心诉求,补充所有澄清得到的限定条件。 127 2. 语句通顺自然,是一个完整的问题,不要分点。 128 3. 只输出最终的查询语句,不要输出其他任何内容。 129 130 原始问题:{original_question} 131 澄清问答记录: 132 {clarify_history} 133 134 最终查询:""" 135 ) 136 137 def generate_clarify_questions(question: str, max_num: int = 3) -> List[str]: 138 chain = CLARIFY_GENERATE_PROMPT | llm | line_parser 139 result = chain.invoke({"question": question, "max_clarify_num": max_num}) 140 return [q.strip() for q in result if q.strip()] 141 142 def integrate_final_query(original_question: str, clarify_history: List[Tuple[str, str]]) -> str: 143 history_text = "\n".join([f"问:{q}\n答:{a}" for q, a in clarify_history]) 144 chain = QUERY_INTEGRATE_PROMPT | llm 145 return chain.invoke({ 146 "original_question": original_question, 147 "clarify_history": history_text 148 }).content.strip() 149 150 def run_clarification( 151 original_question: str, 152 max_rounds: int = 2, 153 questions_per_round: int = 3 154 ) -> str: 155 current_question = original_question 156 clarify_history = [] 157 158 for round_idx in range(max_rounds): 159 clarify_questions = generate_clarify_questions(current_question, questions_per_round) 160 if not clarify_questions: 161 logger.info(f"\n[澄清阶段] 第 {round_idx+1} 轮:问题已足够明确,终止澄清") 162 break 163 164 logger.info(f"\n[澄清阶段] 第 {round_idx+1} 轮澄清问题:") 165 for i, q in enumerate(clarify_questions, 1): 166 logger.info(f" {i}. {q}") 167 168 print("\n请依次回答以下问题:") 169 user_answers = [] 170 for q in clarify_questions: 171 ans = input(f"Q: {q}\nA: ") 172 user_answers.append(ans) 173 174 for q, a in zip(clarify_questions, user_answers): 175 clarify_history.append((q, a)) 176 177 current_question = integrate_final_query(original_question, clarify_history) 178 logger.info(f"\n[澄清阶段] 本轮整合后查询:{current_question}") 179 180 return current_question 181 182 # ===================== 5. 模块二:复合问题拆分 ===================== 183 def split_complex_question(question: str) -> List[str]: 184 prompt = f"""请将用户的问题拆分成若干个独立的单一子问题,每个子问题只包含一个疑问点。 185 要求: 186 1. 每个子问题单独占一行 187 2. 不要编号、不要多余解释 188 3. 保留原问题的所有限定条件和语义 189 190 用户问题:{question}""" 191 response = llm.invoke(prompt) 192 sub_questions = [line.strip() for line in response.content.splitlines() if line.strip()] 193 return sub_questions 194 195 # ===================== 6. 模块三:HyDE 假设文档召回 ===================== 196 HYDE_PROMPT = ChatPromptTemplate.from_template("""请撰写一段与以下问题相关的游戏设定内容,要求表述风格和游戏官方设定文档一致,语义连贯,包含具体细节。 197 问题:{question} 198 内容:""") 199 200 generate_hyde_doc = HYDE_PROMPT | llm | str_parser 201 202 def hyde_retrieve(query: str, vectorstore, top_k: int = 2): 203 """HyDE单路检索,传入指定向量库,生成失败自动跳过""" 204 try: 205 hypo_doc = generate_hyde_doc.invoke({"question": query}) 206 if not hypo_doc or len(hypo_doc.strip()) < 20: 207 logger.warning(f" [HyDE] 生成内容无效,跳过该路召回") 208 return [] 209 210 logger.info(f" [HyDE] 生成假设文档:{hypo_doc[:80]}...") 211 hyde_retriever = vectorstore.as_retriever(search_kwargs={"k": top_k}) 212 docs = hyde_retriever.invoke(hypo_doc) 213 logger.info(f" [HyDE] 召回文档数:{len(docs)}") 214 return docs 215 except Exception as e: 216 logger.warning(f" [HyDE] 执行异常,跳过:{str(e)}") 217 return [] 218 219 # ===================== 7. 模块四:多路召回 + 全局去重 ===================== 220 def single_query_retrieve(query: str, base_retriever, vectorstore, enable_hyde: bool = False): 221 """单个查询执行多路召回:MultiQuery + 可选HyDE""" 222 # 第1路:MultiQuery多角度检索 223 docs = base_retriever.invoke(query) 224 225 # 第2路:HyDE假设文档检索 226 if enable_hyde: 227 hyde_docs = hyde_retrieve(query, vectorstore) 228 docs.extend(hyde_docs) 229 230 return docs 231 232 def batch_retrieve(sub_questions: List[str], base_retriever, vectorstore, enable_hyde: bool = False): 233 """批量子问题检索 + 全局去重""" 234 all_docs = [] 235 for idx, q in enumerate(sub_questions, 1): 236 logger.info(f"\n[检索阶段] 正在处理子问题 {idx}/{len(sub_questions)}:{q}") 237 docs = single_query_retrieve(q, base_retriever, vectorstore, enable_hyde) 238 all_docs.extend(docs) 239 240 # 按文档内容全局去重 241 unique_docs = [] 242 seen_contents = set() 243 for doc in all_docs: 244 if doc.page_content not in seen_contents: 245 seen_contents.add(doc.page_content) 246 unique_docs.append(doc) 247 return unique_docs 248 249 # ===================== 8. 主流水线:全模块串联 ===================== 250 def retrieval_pipeline( 251 original_question: str, 252 enable_routing: bool = False, 253 enable_clarification: bool = True, 254 enable_split: bool = True, 255 enable_hyde: bool = False, 256 max_clarify_rounds: int = 2, 257 questions_per_round: int = 3 258 ): 259 """ 260 完整检索前处理流水线 261 :param original_question: 用户原始问题 262 :param enable_routing: 是否开启逻辑路由(多数据源分发) 263 :param enable_clarification: 是否开启查询澄清 264 :param enable_split: 是否开启复合问题拆分 265 :param enable_hyde: 是否开启HyDE假设文档召回 266 :param max_clarify_rounds: 最大澄清轮数 267 :param questions_per_round: 单轮澄清问题数 268 :return: 最终去重后的文档列表 269 """ 270 logger.info("=" * 60) 271 logger.info(f"原始问题:{original_question}") 272 logger.info(f"功能开关:路由={enable_routing} | 澄清={enable_clarification} | 拆分={enable_split} | HyDE={enable_hyde}") 273 logger.info("=" * 60) 274 275 # 第0步:逻辑路由,选择对应数据源的检索器与向量库 276 target_datasource = DEFAULT_DATASOURCE 277 if enable_routing: 278 target_datasource = run_routing(original_question) 279 280 # 获取当前数据源对应的检索组件 281 current_retriever = multiquery_retriever_map.get(target_datasource, multiquery_retriever_map[DEFAULT_DATASOURCE]) 282 current_vectorstore = vectorstore_map.get(target_datasource, vectorstore_map[DEFAULT_DATASOURCE]) 283 284 # 第一步:查询澄清 285 current_question = original_question 286 if enable_clarification: 287 current_question = run_clarification( 288 current_question, 289 max_rounds=max_clarify_rounds, 290 questions_per_round=questions_per_round 291 ) 292 logger.info(f"\n[澄清完成] 最终明确问题:{current_question}") 293 294 # 第二步:复合问题拆分 295 sub_questions = [current_question] 296 if enable_split: 297 sub_questions = split_complex_question(current_question) 298 logger.info(f"\n[拆分完成] 拆分为 {len(sub_questions)} 个子问题:") 299 for i, q in enumerate(sub_questions, 1): 300 logger.info(f" {i}. {q}") 301 302 # 第三步:多路检索 + 全局去重 303 final_docs = batch_retrieve(sub_questions, current_retriever, current_vectorstore, enable_hyde) 304 305 logger.info("\n" + "=" * 60) 306 logger.info(f"流水线执行完成,共召回 {len(final_docs)} 条去重文档") 307 logger.info("=" * 60) 308 return final_docs 309 310 # ===================== 运行示例 ===================== 311 if __name__ == "__main__": 312 # 测试1:游戏设定类问题(开启路由+HyDE,关闭澄清方便演示) 313 print("\n" + "#"*60) 314 print("测试1:游戏设定类问题") 315 print("#"*60) 316 docs = retrieval_pipeline( 317 original_question="普陀山的BOSS怎么打?", 318 enable_routing=True, 319 enable_clarification=False, 320 enable_split=False, 321 enable_hyde=True 322 ) 323 print("最终检索结果:") 324 for doc in docs: 325 print(f"- {doc.page_content}") 326 327 # 测试2:剧情类问题 328 print("\n" + "#"*60) 329 print("测试2:剧情类问题") 330 print("#"*60) 331 docs = retrieval_pipeline( 332 original_question="孙悟空和黑熊精认识吗?", 333 enable_routing=True, 334 enable_clarification=False, 335 enable_split=False, 336 enable_hyde=False 337 ) 338 print("最终检索结果:") 339 for doc in docs: 340 print(f"- {doc.page_content}") 341 342 # 测试3:客服类问题 343 print("\n" + "#"*60) 344 print("测试3:客服类问题") 345 print("#"*60) 346 docs = retrieval_pipeline( 347 original_question="充了钱没到账要找谁?", 348 enable_routing=True, 349 enable_clarification=False, 350 enable_split=False, 351 enable_hyde=False 352 ) 353 print("最终检索结果:") 354 for doc in docs: 355 print(f"- {doc.page_content}")