Python搭建RetrievalQA链:如何限制仅返回单次数据源匹配结果
解决RetrievalQA链非预期额外推理的方法
以下是几种限制RetrievalQA仅执行单次迭代、避免多余内容的具体方案:
1. 严格约束Prompt模板,禁止额外推理
核心是在Prompt中明确指令,让LLM仅基于检索到的数据源输出指定格式JSON,杜绝额外生成。
示例代码:
from langchain.prompts import PromptTemplate from langchain.chains import RetrievalQA from langchain.llms import OpenAI # 自定义严格约束的Prompt prompt_template = """仅基于以下检索到的数据源,输出符合要求的JSON格式结果,不得添加任何额外解释、推理或无关内容: 数据源:{context} 用户问题:{question} 指定JSON格式示例:{"数据源名称": "xxx", "结果": "xxx"} 输出:""" PROMPT = PromptTemplate( template=prompt_template, input_variables=["context", "question"] ) # 初始化RetrievalQA链,使用自定义Prompt并指定chain_type为"stuff" qa_chain = RetrievalQA.from_chain_type( llm=OpenAI(temperature=0), # 温度设为0减少随机性 chain_type="stuff", retriever=your_vector_db_retriever, chain_type_kwargs={"prompt": PROMPT}, return_source_documents=False # 不需要返回源文档,避免干扰 ) # 调用链 result = qa_chain.run("食品相关商品总销售额是多少?")
2. 使用结构化输出解析器强制格式
利用结构化输出工具,强制LLM输出符合Schema的JSON,从根源限制多余内容。
示例代码:
from langchain.output_parsers import StructuredOutputParser, ResponseSchema from langchain.prompts import PromptTemplate from langchain.chains import RetrievalQA # 定义输出Schema response_schemas = [ ResponseSchema(name="数据源名称", description="匹配到的数据源名称"), ResponseSchema(name="结果", description="用户问题的答案内容") ] output_parser = StructuredOutputParser.from_response_schemas(response_schemas) format_instructions = output_parser.get_format_instructions() # 构建带格式约束的Prompt prompt = PromptTemplate( template="仅基于以下数据源回答问题,严格按照指定格式输出JSON,不得添加其他内容:\n{context}\n\n问题:{question}\n\n{format_instructions}", input_variables=["context", "question"], partial_variables={"format_instructions": format_instructions} ) # 初始化链 qa_chain = RetrievalQA.from_chain_type( llm=OpenAI(temperature=0), chain_type="stuff", retriever=your_vector_db_retriever, chain_type_kwargs={"prompt": prompt}, return_source_documents=False ) # 调用并解析结果 result = qa_chain.run("食品相关商品总销售额是多少?") parsed_result = output_parser.parse(result)
3. 拆分链流程,手动控制检索与输出
跳过RetrievalQA的默认封装,手动拆分“检索数据源”和“生成JSON”两步,完全控制迭代次数。
示例代码:
# 第一步:从向量库检索匹配的数据源 retrieved_docs = your_vector_db_retriever.get_relevant_documents("食品相关商品总销售额是多少?") context = "\n".join([doc.page_content for doc in retrieved_docs]) source_name = retrieved_docs[0].metadata["source"] # 假设元数据中存数据源名称 # 第二步:调用LLM生成指定格式JSON llm = OpenAI(temperature=0) prompt = f"""仅基于以下数据源生成JSON结果,不得添加额外内容: 数据源:{context} 数据源名称:{source_name} 问题:食品相关商品总销售额是多少? 输出格式:{{"数据源名称": "{source_name}", "结果": "xxx"}}""" result = llm(prompt)
4. 限制LLM输出参数
通过设置max_tokens和stop序列,强制LLM在输出完JSON后停止,避免多余推理内容。
示例代码:
from langchain.chains import RetrievalQA from langchain.llms import OpenAI qa_chain = RetrievalQA.from_chain_type( llm=OpenAI( temperature=0, max_tokens=200, # 根据预期JSON长度设置合适值 stop=["\n"] # 若JSON是单行,设置换行符为停止标志 ), chain_type="stuff", retriever=your_vector_db_retriever, chain_type_kwargs={"prompt": PROMPT} # 配合之前的约束Prompt )
内容的提问来源于stack exchange,提问作者Frank Pinto
相关产品推荐
相关产品推荐

