如何正确解析load_qa_chain输出?解决JSONDecodeError问题
解决LangChain load_qa_chain结合PydanticOutputParser的JSON解析错误问题
问题核心在于load_qa_chain的默认Prompt不会强制LLM输出严格符合Pydantic要求的JSON结构,返回的文本可能包含多余说明、格式不规范,导致parser.parse()触发JSONDecodeError。以下是具体修复方案及代码示例:
关键修复步骤
- 明确Pydantic模型定义:确保输出结构清晰,字段描述准确,让LLM理解要生成的内容格式。
- 自定义Prompt模板:在Prompt中强制要求LLM仅输出指定格式的JSON,禁止添加任何额外自然语言内容,并注入Pydantic的格式说明。
- 替换load_qa_chain的默认Prompt:将自定义Prompt传入chain,确保生成的输出符合解析要求。
完整代码示例
1. 定义Pydantic模型与解析器
from pydantic import BaseModel, Field from langchain.output_parsers import PydanticOutputParser # 定义期望的双结果结构 class DualResult(BaseModel): result1: str = Field(description="第一个问题的答案内容") result2: str = Field(description="第二个问题的答案内容") # 初始化解析器 parser = PydanticOutputParser(pydantic_object=DualResult)
2. 自定义Prompt并初始化load_qa_chain
from langchain.prompts import PromptTemplate from langchain.chains.question_answering import load_qa_chain from langchain.llms import OpenAI # 构建严格约束输出格式的Prompt prompt_template = """基于以下上下文回答问题,必须严格按照指定的JSON格式输出,**禁止添加任何额外解释或文本**: {context} 问题: {question} {format_instructions} """ PROMPT = PromptTemplate( template=prompt_template, input_variables=["context", "question"], # 注入Pydantic的格式说明,让LLM明确输出结构 partial_variables={"format_instructions": parser.get_format_instructions()} ) # 初始化LLM与chain(这里以OpenAI为例,可替换为其他LLM) llm = OpenAI(temperature=0) # 使用stuff类型chain(适合短文档场景) chain = load_qa_chain(llm, chain_type="stuff", prompt=PROMPT)
3. 调用chain并解析结果
# 假设docs是你检索到的上下文文档列表 docs = [...] # 你的问题(可以包含需要返回两个结果的需求) question = "请根据上下文,分别回答XXX和YYY两个问题" # 执行chain result = chain.run(input_documents=docs, question=question) # 解析结果 parsed_result = parser.parse(result) # 访问解析后的结果 print(parsed_result.result1) print(parsed_result.result2)
长文档场景适配(map_reduce类型chain)
如果处理长文档需要使用map_reduce类型chain,需要分别定义map阶段和combine阶段的Prompt,确保每一步输出都符合格式要求:
# map阶段Prompt:处理单段文档生成JSON片段 map_template = """基于以下文档片段回答问题,输出符合要求的JSON: {context} 问题: {question} {format_instructions} """ MAP_PROMPT = PromptTemplate( template=map_template, input_variables=["context", "question"], partial_variables={"format_instructions": parser.get_format_instructions()} ) # combine阶段Prompt:整合所有片段结果生成最终JSON combine_template = """整合以下所有结果,生成最终的符合要求的JSON,**禁止添加额外内容**: {summaries} {format_instructions} """ COMBINE_PROMPT = PromptTemplate( template=combine_template, input_variables=["summaries"], partial_variables={"format_instructions": parser.get_format_instructions()} ) # 初始化map_reduce类型chain chain = load_qa_chain(llm, chain_type="map_reduce", map_prompt=MAP_PROMPT, combine_prompt=COMBINE_PROMPT)
内容的提问来源于stack exchange,提问作者rich
相关产品推荐
相关产品推荐

