如何在LangChain的RetrievalQA中添加记忆与多输入自定义提示词
在LangChain的RetrievalQA中集成对话记忆与多输入自定义提示词
要解决RetrievalQA无法同时支持多输入自定义提示词和对话记忆的问题,需要调整链的结构配置,确保额外参数(客户信息)和对话历史能正确传递到提示词中。以下是具体解决方案和修正后的代码:
关键问题分析
原代码存在两个核心问题:
- RetrievalQA的
chain_type_kwargs中直接传入memory的方式不正确,RetrievalQA本身不支持这种配置方式,对话记忆需要通过正确的链绑定逻辑传递。 - 自定义提示词中的多输入参数(
Customer_Name、Customer_State、Customer_Gender)需要被显式传递到QA链的调用上下文里,确保提示词能正确渲染。
修正后的实现代码
import openai import numpy as np import pandas as pd import os import json from langchain.embeddings.sentence_transformer import SentenceTransformerEmbeddings from langchain.vectorstores import Chroma from langchain.text_splitter import RecursiveCharacterTextSplitter from langchain.llms import OpenAI from langchain.chains import RetrievalQA from langchain.chains.question_answering import load_qa_chain from langchain.document_loaders import UnstructuredFileLoader from langchain.prompts import PromptTemplate from langchain.memory import ConversationBufferMemory from langchain.docstore.document import Document # 加载PDF文档 loader = UnstructuredFileLoader("../Test.pdf", mode="elements") documents = loader.load() # 加载客户信息 with open('Customer_profile.json', 'r') as openfile: json_object = json.load(openfile) cName = json_object['Customer_Name'] cState = json_object['Customer_State'] cGen = json_object['Customer_Gender'] # 定义带多输入的自定义提示词 prompt_template = """You are a Chat customer support agent. Address the customer as Dear Mr. or Miss. depending on customer's gender followed by Customer's First Name. Use the following customer related information (delimited by <cp></cp>) context (delimited by <ctx></ctx>) and the chat history (delimited by <hs></hs>) to answer the question at the end: If you don't know the answer, just say that you don't know, don't try to make up an answer. Below are the details of the customer: <cp> Customer's Name: {Customer_Name} Customer's Resident State: {Customer_State} Customer's Gender: {Customer_Gender} </cp> <ctx> {context} </ctx> <hs> {history} </hs> Question: {query} Answer: """ PROMPT = PromptTemplate( template=prompt_template, input_variables=["history", "context", "query", "Customer_Name", "Customer_State", "Customer_Gender"] ) # 文本分块与向量存储构建 text_splitter = RecursiveCharacterTextSplitter(chunk_size=1000, chunk_overlap=0) texts = text_splitter.split_documents(documents) embeddings = SentenceTransformerEmbeddings(model_name="all-MiniLM-L6-v2") vectorDB = Chroma.from_documents(texts, embeddings) # 配置对话记忆:只存储对话历史,输入key匹配query memory = ConversationBufferMemory( memory_key="history", input_key="query", return_messages=False # 直接返回字符串格式的历史,适配提示词 ) # 加载自定义QA链,使用stuff模式 qa_chain = load_qa_chain( llm=OpenAI(), chain_type="stuff", prompt=PROMPT, verbose=True ) # 构建RetrievalQA,将自定义链传入,并绑定记忆 qa = RetrievalQA( combine_documents_chain=qa_chain, retriever=vectorDB.as_retriever(), verbose=True, memory=memory, return_source_documents=False ) # 调用链,传入所有必要参数 response = qa({ "query": "who's the client's friend?", "Customer_Gender": cGen, "Customer_State": cState, "Customer_Name": cName }) print(response['result']) # 后续对话示例,记忆会自动携带历史 response2 = qa({ "query": "What's his contact information?", "Customer_Gender": cGen, "Customer_State": cState, "Customer_Name": cName }) print(response2['result'])
核心修改说明
- 对话记忆的正确绑定:将
memory直接作为RetrievalQA的初始化参数传入,而非放在chain_type_kwargs中,确保RetrievalQA能正确管理对话历史的存储与传递;设置return_messages=False,让记忆返回纯文本格式的对话历史,适配自定义提示词的渲染需求。 - 多输入参数的传递:调用链时必须显式传入所有自定义提示词中定义的输入变量,这些参数会被自动传递到提示词中完成渲染。
- 自定义QA链的加载:使用
load_qa_chain加载自定义提示词生成符合需求的QA链,再将其传入RetrievalQA的combine_documents_chain参数,替代默认链类型配置,实现完全自定义的提示词逻辑。
内容的提问来源于stack exchange,提问作者Jason
相关产品推荐
相关产品推荐

