如何为Gradio ChatInterface实现用户独立的会话状态?
解决Gradio ChatInterface用户会话历史共享问题
问题根源在于你全局初始化了chain实例,所有用户共用同一个ConversationBufferMemory,导致聊天历史跨用户共享。要实现用户独立会话,需要为每个会话创建独立的chain实例,并通过Gradio的会话状态管理这些实例。
修改方案
1. 移除全局chain实例
不在全局作用域创建chain = load_chain(),改为在会话开始时为每个用户单独初始化。
2. 利用Gradio会话状态存储用户专属chain
通过gr.State组件为每个用户会话维护独立的chain实例,确保会话间状态完全隔离。
3. 调整predict函数逻辑
修改predict函数,接收会话状态中的chain,若不存在则创建新实例;使用该chain处理当前用户对话,保证历史仅属于当前会话。
修改后的完整代码
import os from typing import Optional, Tuple import gradio as gr from langchain.chains import ConversationChain, ConversationalRetrievalChain from langchain.document_loaders import PyPDFLoader from langchain.text_splitter import RecursiveCharacterTextSplitter from langchain.embeddings import HuggingFaceEmbeddings from langchain.vectorstores import Chroma from langchain.memory import ConversationBufferMemory from langchain.chat_models import ChatOpenAI from langchain.prompts import PromptTemplate from langchain.retrievers import ContextualCompressionRetriever from langchain.retrievers.document_compressors import LLMChainExtractor from langchain.llms import OpenAI from langchain.schema import AIMessage, HumanMessage def load_chain(llm_name="gpt-3.5-turbo"): """Logic for loading the chain you want to use should go here.""" # define embedding embedding = HuggingFaceEmbeddings( model_name="sentence-transformers/all-mpnet-base-v2", model_kwargs={'device': 'cpu'}, encode_kwargs={'normalize_embeddings': False} ) # create vector database from data persist_directory = 'docs/chroma/' vectordb = Chroma(persist_directory=persist_directory, embedding_function=embedding) # Wrap our vectorstore llm = OpenAI(temperature=0) compressor = LLMChainExtractor.from_llm(llm) # define retriever compression_retriever = ContextualCompressionRetriever( base_compressor=compressor, base_retriever=vectordb.as_retriever(search_type = "mmr") ) # Build prompt template = """ Use the following pieces of context 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. If there are any assumptions or requirements for the answer to apply, please include them in your response. {context} Question: {question} Helpful Answer:""" QA_CHAIN_PROMPT = PromptTemplate(input_variables=["context", "question"],template=template,) # define memory memory = ConversationBufferMemory(memory_key="chat_history", return_messages=True) # create a chatbot chain chain = ConversationalRetrievalChain.from_llm( llm=ChatOpenAI(model_name=llm_name, temperature=0), memory=memory, retriever=compression_retriever, combine_docs_chain_kwargs={"prompt": QA_CHAIN_PROMPT} ) return chain def predict(message, history, chain): # 会话无chain实例时,初始化新的专属实例 if chain is None: chain = load_chain() # 用当前用户的chain处理对话,自动维护专属历史 gpt_response = chain({"question": message}) return gpt_response['answer'], chain block = gr.Blocks() with block: # 用State存储每个用户的chain实例,初始值为None chain_state = gr.State(value=None) chatbot = gr.ChatInterface( fn=predict, title="Chatbot", # 将chain_state作为额外输入,接收更新后的chain作为输出 additional_inputs=[chain_state], additional_outputs=[chain_state] ) block.launch()
关键修改说明
- 会话状态隔离:
gr.State会为每个用户会话创建独立的chain实例,不同用户的聊天历史完全分离。 - 延迟初始化:仅在用户发起第一次对话时创建chain,避免不必要的资源消耗。
- 状态持续更新:
predict函数返回更新后的chain,确保会话状态始终保存当前用户的聊天历史。
部署修改后的代码到HuggingFace Spaces后,每个用户(无论使用哪个浏览器或设备)都会拥有独立的聊天会话,关闭页面或会话结束后,该用户的历史会自动销毁。
内容的提问来源于stack exchange,提问作者Angela Y
相关产品推荐
相关产品推荐

