Langchain中RetrievalQA切换map_reduce链类型时Prompt配置报错求助
解决Langchain RetrievalQA使用map_reduce链时的ValidationError问题
问题概述
使用Langchain的ParentDocumentRetriever结合对话记忆构建RAG模型,默认chain_type="stuff"运行正常,但切换为map_reduce链类型时,触发以下错误:
ValidationError: 1 validation error for RefineDocumentsChain
prompt
extra fields not permitted (type=value_error.extra)
错误原因
map_reduce链类型不接受单个prompt参数,它需要分别配置**map_prompt(处理每段检索到的文档,生成初步回答)和combine_prompt**(合并所有初步回答,生成最终结果)。此外,原代码将对话记忆放在chain_type_kwargs中,不符合map_reduce链的参数结构要求。
解决方案
1. 拆分构建map和combine阶段的Prompt
根据交互式RAG需求,分别定义两个Prompt模板:
- Map阶段模板:针对单份文档片段生成初步回答,无需对话历史(历史将在合并阶段引入)
- Combine阶段模板:汇总所有初步回答,结合对话历史生成最终的德语回答
2. 调整链参数配置
将对话记忆从chain_type_kwargs移至RetrievalQA.from_chain_type的memory参数中,同时在chain_type_kwargs中传入map_prompt和combine_prompt。
修改后的完整代码
from langchain.chains import RetrievalQA from langchain.memory import ConversationSummaryMemory from langchain.prompts import PromptTemplate from langchain.document_loaders import PyPDFLoader, DirectoryLoader from langchain.text_splitter import RecursiveCharacterTextSplitter from langchain.vectorstores import Chroma from langchain.storage import InMemoryStore from chromadb.errors import InvalidDimensionException # 加载文档 loader = DirectoryLoader("MY_PATH_TO_PDF_FILES", glob='*.pdf', loader_cls=PyPDFLoader) documents = loader.load() # 定义文档拆分器 parent_splitter = RecursiveCharacterTextSplitter(chunk_size=2000, chunk_overlap=400) child_splitter = RecursiveCharacterTextSplitter(chunk_size=400) # 初始化向量库 try: vectorstore = Chroma(collection_name="split_parents", embedding_function=bge_embeddings, persist_directory="chroma_db") except InvalidDimensionException: Chroma().delete_collection() vectorstore = Chroma(collection_name="split_parents", embedding_function=bge_embeddings, persist_directory="chroma_db") # 初始化父文档存储 store = InMemoryStore() # 初始化ParentDocumentRetriever big_chunks_retriever = ParentDocumentRetriever( vectorstore=vectorstore, docstore=store, child_splitter=child_splitter, parent_splitter=parent_splitter, ) big_chunks_retriever.add_documents(documents) # -------------------------- # 定义map和combine阶段的Prompt # -------------------------- # Map阶段:处理单个文档片段,生成初步回答 map_template = """ 使用以下上下文信息回答问题,仅用德语作答! 如果不知道答案,回答"Leider habe ich keine Informationen." ------ 上下文: {context} ------ 问题:{question} 初步回答: """ map_prompt = PromptTemplate(template=map_template, input_variables=["context", "question"]) # Combine阶段:合并所有初步回答,结合对话历史生成最终答案 combine_template = """ 根据以下所有初步回答和对话历史,回答用户的问题,仅用德语作答! 如果没有足够信息或不知道答案,回答"Leider habe ich keine Informationen." ------ 对话历史: {chat_history} ------ 所有初步回答: {summaries} ------ 问题:{question} 最终回答: """ combine_prompt = PromptTemplate(template=combine_template, input_variables=["summaries", "chat_history", "question"]) # 初始化对话记忆 memory = ConversationSummaryMemory( llm=build_llm(), memory_key="chat_history", # 与combine_prompt中的变量名对应 input_key="question", return_messages=True ) # 配置chain_type_kwargs chain_type_kwargs = { "verbose": True, "map_prompt": map_prompt, "combine_prompt": combine_prompt } # 初始化RetrievalQA(map_reduce链类型) qa_chain = RetrievalQA.from_chain_type( llm=build_llm(), chain_type="map_reduce", return_source_documents=True, chain_type_kwargs=chain_type_kwargs, retriever=big_chunks_retriever, memory=memory, # 对话记忆直接传入RetrievalQA verbose=True ) # 测试查询 query = "Hi, I am Max, can you help me??" result = qa_chain(query) print(result["result"])
关键修改点说明
- Prompt拆分:将原单一模板拆分为
map_prompt和combine_prompt,分别适配map和combine阶段的任务需求 - 记忆参数调整:将
ConversationSummaryMemory从chain_type_kwargs移至RetrievalQA.from_chain_type的memory参数中,确保对话历史能正确传入combine阶段 - 变量名对齐:
memory_key设置为chat_history,与combine_prompt中的变量名保持一致,避免参数匹配错误 - 语言规范:将未知答案的回复改为德语,符合用户要求
内容的提问来源于stack exchange,提问作者Maxl Gemeinderat
相关产品推荐
相关产品推荐

