RAG项目中ConversationalRetrievalChain对话记忆失效问题求助
问题排查与修复
核心问题分析
你的代码存在几个关键错误,导致对话历史无法被模型识别:
- 每次请求重建对话链与记忆:每次点击按钮调用
get_response时,都会重新创建ConversationBufferMemory和ConversationalRetrievalChain,之前的对话记忆无法被新链加载。 - 记忆格式不匹配:
ConversationBufferMemory设置了return_messages=True,要求对话历史是BaseMessage对象列表,但你存储的是字符串元组,链无法正确解析历史内容。 - 手动传参与内存组件冲突:同时使用
memory参数和手动传入chat_history到链的调用中,导致内存组件无法正常工作。 - 冗余的Prompt模板:
PROMPT_TEMPLATE给用户问题额外添加了"Question:"前缀,既干扰模型理解,也让历史记录格式混乱。
修复后的完整代码
import time from typing import Dict, Optional import streamlit as st from langchain_community.vectorstores.pinecone import Pinecone as LangChainPinecone from langchain_openai.embeddings import OpenAIEmbeddings from langchain_openai.llms import OpenAI from langchain.memory import ConversationBufferMemory from pinecone import Pinecone from langchain.chains import ConversationalRetrievalChain import os INDEX_NAME = "index-name" NUM_RETRIEVED_DOCS = 5 TEMPERATURE = 0.3 CONVERSATION_MEMORY_SIZE = 5 def initialize_pinecone_client(api_key: str) -> Pinecone: return Pinecone(api_key=api_key) def initialize_session_state(): if "chat_history" not in st.session_state: st.session_state.chat_history = [] # 将对话链和内存存入Session State,避免每次请求重建 if "qa_chain" not in st.session_state: pc_client = initialize_pinecone_client(api_key=st.secrets["pinecone_api_key"]) embedding_model = OpenAIEmbeddings(model="text-embedding-3-large", openai_api_key=st.secrets["openai_api_key"]) index = pc_client.Index(INDEX_NAME) time.sleep(1) index.describe_index_stats() # 初始化内存,关闭return_messages以匹配元组格式 conversation_memory = ConversationBufferMemory( memory_key="chat_history", return_messages=False, k=CONVERSATION_MEMORY_SIZE ) vectorstore = LangChainPinecone(index=index, embedding=embedding_model, text_key="context") llm = OpenAI(temperature=TEMPERATURE, openai_api_key=st.secrets["openai_api_key"]) retriever = vectorstore.as_retriever(search_type="similarity", search_kwargs={"k": NUM_RETRIEVED_DOCS}) st.session_state.qa_chain = ConversationalRetrievalChain.from_llm( llm=llm, retriever=retriever, memory=conversation_memory, return_source_documents=True ) def get_response(user_query: str) -> Dict[str, Optional[str]]: try: initialize_session_state() # 复用Session State中已初始化的对话链 result = st.session_state.qa_chain({"question": user_query}) # 同步更新前端展示用的对话历史 st.session_state.chat_history.append((user_query, result['answer'])) return {'result': result, 'chat_history': st.session_state.chat_history} except Exception as e: print(f"An error occurred: {e}") return {"answer": "处理请求时出错,请稍后重试。", "sources": None} user_query = st.text_area("Enter your question here:") if st.button("Get Response"): if user_query: response = get_response(user_query) st.write("Answer:", response['result'].get("answer", "未找到答案。")) sources = response['result'].get("source_documents") if sources: for idx, doc in enumerate(sources): st.write(f"Source {idx+1}:", os.path.basename(doc.metadata.get('source', 'Unknown'))) else: st.write("未引用特定来源。") st.write("Chat History:") for i, (question, answer) in enumerate(response['chat_history']): st.write(f"Turn {i + 1}:") st.write(f"Q: {question}") st.write(f"A: {answer}") st.write("---") else: st.error("请输入问题。")
关键修复点说明
- 会话状态持久化对话链:将
qa_chain和conversation_memory的初始化逻辑移到initialize_session_state中,并存储在st.session_state,确保每次请求复用同一个链和记忆组件,保留对话历史。 - 统一记忆格式:将
ConversationBufferMemory的return_messages设为False,匹配你存储的字符串元组格式,让链能正确读取历史内容。 - 移除冗余参数与模板:删除手动传入
chat_history的代码,去掉PROMPT_TEMPLATE,直接使用用户原始输入的问题,避免格式干扰。 - 同步展示用历史列表:对话链的内存会自动管理历史,这里同步更新
st.session_state.chat_history用于前端展示。
内容的提问来源于stack exchange,提问作者kaanyvz
相关产品推荐
相关产品推荐

