检索前处理:查询构建、查询翻译、查询路由

自己搞的简单代码,便于理解

  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}")