You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何为load_qa_chain添加记忆或实现带多输入自定义Prompt的ConversationalRetrievalChain

解决方案:LangChain问答功能的多输入Prompt与对话记忆兼容问题

方案一:为load_qa_chain添加对话记忆

load_qa_chain本身没有内置对话记忆支持,但可以手动结合记忆组件实现。核心思路是:每次调用时从记忆中取出历史对话,传入自定义Prompt,调用完成后将当前问答对存入记忆。

实现步骤与代码示例

  1. 定义包含历史对话和多输入参数的自定义Prompt
from langchain.prompts import PromptTemplate
from langchain.chains.question_answering import load_qa_chain
from langchain.memory import ConversationBufferMemory
from langchain.llms import OpenAI

# 自定义Prompt,包含历史对话、用户问题、额外参数(比如用户角色)
prompt_template = """
历史对话:
{chat_history}

用户角色:{user_role}
用户问题:{question}
上下文:{context}

请基于上下文和历史对话,以符合用户角色的语气回答问题:
"""
PROMPT = PromptTemplate(
    template=prompt_template,
    input_variables=["chat_history", "user_role", "question", "context"]
)
  1. 初始化记忆组件和QA链
# 初始化记忆,指定记忆存储的键为"chat_history"
memory = ConversationBufferMemory(memory_key="chat_history")

# 初始化LLM和QA链
llm = OpenAI(temperature=0)
qa_chain = load_qa_chain(llm, chain_type="stuff", prompt=PROMPT)
  1. 封装调用逻辑,整合记忆读写
def chat_with_memory(user_question, user_role, docs):
    # 从记忆中获取历史对话
    chat_history = memory.load_memory_variables({})["chat_history"]
    
    # 调用QA链,传入所有参数
    result = qa_chain({
        "input_documents": docs,
        "question": user_question,
        "user_role": user_role,
        "chat_history": chat_history
    })
    
    # 构造当前问答对,存入记忆
    current_chat = f"用户:{user_question}\n助手:{result['output_text']}"
    memory.save_context({"input": user_question}, {"output": result["output_text"]})
    
    return result["output_text"]

方案二:让ConversationalRetrievalChain支持多输入自定义Prompt

ConversationalRetrievalChain默认的prompt只支持有限参数,但可以通过自定义Chain或者扩展其输入来实现多参数支持,核心是让Chain能接收并传递额外参数到Prompt中。

实现步骤与代码示例

  1. 定义包含多输入参数的自定义Prompt
from langchain.prompts import PromptTemplate
from langchain.chains import ConversationalRetrievalChain
from langchain.memory import ConversationBufferMemory
from langchain.llms import OpenAI
from langchain.chains.question_answering import load_qa_chain

# 自定义多输入Prompt,包含历史对话、用户问题、额外参数(比如领域)
qa_prompt = PromptTemplate(
    template="""
历史对话:{chat_history}
领域:{domain}
用户问题:{question}
上下文:{context}

请结合领域知识、历史对话和上下文回答问题:
""",
    input_variables=["chat_history", "domain", "question", "context"]
)
  1. 自定义QA链,传递多参数
# 初始化LLM和带自定义Prompt的QA链
llm = OpenAI(temperature=0)
qa_chain = load_qa_chain(llm, chain_type="stuff", prompt=qa_prompt)

# 初始化记忆
memory = ConversationBufferMemory(memory_key="chat_history", return_messages=True)

# 自定义ConversationalRetrievalChain的调用逻辑,支持额外参数
class CustomConversationalRetrievalChain(ConversationalRetrievalChain):
    def __init__(self, **kwargs):
        super().__init__(**kwargs)
    
    def _call(self, inputs):
        # 提取额外参数(比如domain)
        domain = inputs.get("domain")
        # 获取历史对话
        chat_history = inputs.get(self.memory_key, "")
        # 获取用户问题
        question = inputs["question"]
        
        # 检索相关文档
        docs = self.retriever.get_relevant_documents(question)
        
        # 调用QA链,传入所有参数
        result = self.combine_docs_chain.run(
            input_documents=docs,
            question=question,
            chat_history=chat_history,
            domain=domain
        )
        
        # 更新记忆
        self.memory.save_context({"input": question}, {"output": result})
        
        return {"answer": result}
  1. 初始化并使用自定义Chain
from langchain.vectorstores import Chroma
from langchain.embeddings import OpenAIEmbeddings

# 初始化检索器(示例用Chroma)
embeddings = OpenAIEmbeddings()
vectorstore = Chroma(persist_directory="./chroma_db", embedding_function=embeddings)
retriever = vectorstore.as_retriever()

# 初始化自定义Chain
custom_chain = CustomConversationalRetrievalChain(
    retriever=retriever,
    combine_docs_chain=qa_chain,
    memory=memory
)

# 调用示例
response = custom_chain({
    "question": "这个产品的售后政策是什么?",
    "domain": "电商售后"
})
print(response["answer"])

简化替代方案:使用RunnableSequence组合

如果不想自定义Chain,也可以用LangChain的RunnableSequence直接组合记忆、检索和QA链,更灵活地传递多参数:

from langchain.schema.runnable import RunnablePassthrough

# 构造输入组合,传递所有参数
input_combiner = RunnablePassthrough.assign(
    chat_history=lambda x: memory.load_memory_variables({})["chat_history"],
    context=lambda x: retriever.get_relevant_documents(x["question"])
)

# 组合成完整流程
chain = input_combiner | qa_chain

# 调用并更新记忆
def run_chain(question, domain):
    result = chain.invoke({
        "question": question,
        "domain": domain
    })
    memory.save_context({"input": question}, {"output": result})
    return result

内容的提问来源于stack exchange,提问作者Jason

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.15 21:46:00