如何持久化LangChain的ConversationBufferMemory?解决序列化验证错误
ConversationBufferMemory 持久化报错解决方案
问题场景
使用LangChain创建对话链后,尝试通过Pydantic序列化ConversationBufferMemory实现跨会话持久化,代码如下:
llm = ChatOpenAI(temperature=0, openai_api_key=OPENAI_API_KEY, model_name=OPENAI_DEFAULT_MODEL) conversation = ConversationChain(llm=llm, memory=ConversationBufferMemory()) # 尝试保存 saved_dict = conversation.memory.chat_memory.dict() # 尝试加载 cm = ChatMessageHistory(**saved_dict) # 或 cm = ChatMessageHistory.parse_obj(saved_dict)
执行时出现报错:
ValidationError: 6 validation errors for ChatMessageHistory messages -> 0 Can't instantiate abstract class BaseMessage with abstract method type (type=type_error)
报错原因
ChatMessageHistory中的messages是BaseMessage的子类实例(如HumanMessage、AIMessage),直接调用dict()序列化时会丢失子类类型信息。反序列化时,Pydantic无法识别具体要实例化哪个子类,只能尝试创建抽象基类BaseMessage,从而触发错误。
解决方案
方法1:手动处理消息序列化与反序列化
通过显式保存消息类型,加载时根据类型创建对应子类实例:
保存记忆
import json from langchain.memory import ConversationBufferMemory def save_conversation_memory(memory, save_path): # 遍历消息,保存类型、内容及附加参数 serialized_messages = [] for msg in memory.chat_memory.messages: serialized_messages.append({ "type": msg.type, "content": msg.content, "additional_kwargs": msg.additional_kwargs }) with open(save_path, "w", encoding="utf-8") as f: json.dump({"messages": serialized_messages}, f)
加载记忆
import json from langchain.schema import HumanMessage, AIMessage, SystemMessage from langchain.memory import ConversationBufferMemory from langchain.schema import ChatMessageHistory def load_conversation_memory(load_path): # 映射消息类型到对应类 msg_type_map = { "human": HumanMessage, "ai": AIMessage, "system": SystemMessage } with open(load_path, "r", encoding="utf-8") as f: data = json.load(f) messages = [] for msg_data in data["messages"]: msg_cls = msg_type_map.get(msg_data["type"]) if msg_cls: messages.append(msg_cls( content=msg_data["content"], additional_kwargs=msg_data.get("additional_kwargs", {}) )) chat_history = ChatMessageHistory(messages=messages) return ConversationBufferMemory(chat_memory=chat_history)
使用示例
# 保存对话记忆 save_conversation_memory(conversation.memory, "conversation_memory.json") # 加载对话记忆并创建新对话链 loaded_memory = load_conversation_memory("conversation_memory.json") new_conversation = ConversationChain(llm=llm, memory=loaded_memory)
方法2:使用LangChain内置序列化工具(推荐)
LangChain新版本提供了专门的序列化工具,可直接处理消息类型:
from langchain.serialization import loads, dumps # 保存对话记忆 saved_content = dumps(conversation.memory.chat_memory) with open("memory_data.json", "w", encoding="utf-8") as f: f.write(saved_content) # 加载对话记忆 with open("memory_data.json", "r", encoding="utf-8") as f: saved_content = f.read() chat_history = loads(saved_content) # 创建带加载后记忆的对话链 loaded_memory = ConversationBufferMemory(chat_memory=chat_history) new_conversation = ConversationChain(llm=llm, memory=loaded_memory)
内容的提问来源于stack exchange,提问作者Neil C. Obremski
相关产品推荐
相关产品推荐

