如何修改LangChain Memory前缀格式适配Llama2对话模板?
解决Llama2与LangChain对话格式不匹配问题
方法1:自定义对话模板(推荐)
完全按照Llama2要求的格式构建提示模板,从根源上避免格式冲突,示例代码如下:
from langchain.prompts import PromptTemplate from langchain.memory import ConversationBufferMemory, ChatMessageHistory # 定义符合Llama2规范的对话模板 template = """{chat_history} [INST] {input} [/INST]""" prompt = PromptTemplate( input_variables=["chat_history", "input"], template=template ) # 自定义记忆格式化逻辑 class CustomLlama2Memory(ConversationBufferMemory): def format_memory_variables(self, inputs): chat_history = [] for message in self.chat_memory.messages: if message.type == "human": chat_history.append(f"[INST] {message.content} [/INST]") else: chat_history.append(message.content) return {"chat_history": "\n".join(chat_history)} memory = CustomLlama2Memory(memory_key="chat_history") # 将模板与记忆传入对话链使用 # conversation = ConversationChain(prompt=prompt, llm=your_llm, memory=memory)
方法2:重写记忆类的格式化函数
继承默认记忆类,修改对话历史的生成逻辑,移除多余的前缀与冒号:
from langchain.memory import ConversationBufferMemory class Llama2CompatibleMemory(ConversationBufferMemory): def format_memory_variables(self, inputs): formatted_history = [] for msg in self.chat_memory.messages: if msg.type == "human": formatted_history.append(f"[INST] {msg.content} [/INST]") else: formatted_history.append(msg.content) return {"chat_history": "\n".join(formatted_history)} # 初始化自定义记忆 memory = Llama2CompatibleMemory(memory_key="chat_history")
方法3:后处理模型输出
若无法修改模板或记忆逻辑,可在模型返回结果后手动清理多余前缀:
# 假设已初始化对话链conversation response = conversation.predict(input="你的问题内容") # 移除输出中的"AI:"前缀并去除首尾空格 clean_response = response.replace("AI:", "").strip()
内容的提问来源于stack exchange,提问作者AndyLinOuO
相关产品推荐
相关产品推荐

