如何优化meta-llama/Llama-2-13b-chat-hf的Prompt及对话记忆?
解决方案:Llama-2-13b-chat-hf 结合 LangChain 记忆与文档检索的对话逻辑
1. 先对齐Llama-2原生对话格式
Llama-2聊天模型要求严格遵循特定对话格式,这是输出混乱的核心原因之一。格式模板为:
<s>[INST] 系统提示 + 历史对话 + 当前问题 [/INST] 模型回答 </s>
必须将LangChain的记忆内容、用户问题、文档片段(如有)嵌入到该格式的[INST]块中。
2. 整合ConversationSummaryBufferMemory与模型格式
通过自定义PromptTemplate,将记忆模块维护的历史对话摘要+最近完整对话注入到Llama-2的格式中:
from langchain.llms import HuggingFacePipeline from langchain.memory import ConversationSummaryBufferMemory from langchain.prompts import PromptTemplate from langchain.chains import ConversationChain # 初始化Llama-2管道(需确保tokenizer设置pad_token=tokenizer.eos_token) pipe = ... # 替换为你的HuggingFacePipeline初始化代码 # 适配Llama-2的对话Prompt模板 llama2_prompt = PromptTemplate( input_variables=["history", "input", "context"], template="""<s>[INST] 你是专业助手,回答规则如下: 1. 若提供相关文档片段,优先基于片段内容回答,禁止编造信息 2. 若无相关文档片段,结合对话历史与自身知识正常回答 对话历史: {history} 当前问题:{input} {context}[/INST]""" ) # 初始化对话记忆模块 memory = ConversationSummaryBufferMemory( llm=pipe, max_token_limit=1000, memory_key="history", input_key="input" )
3. 实现“文档关联则用片段,否则正常聊天”的分支逻辑
方式一:先检索再分支判断
先检索用户问题相关的文档片段,存在则注入Prompt,否则用空上下文触发普通聊天:
from langchain.vectorstores import FAISS from langchain.embeddings import HuggingFaceEmbeddings # 初始化向量检索库(假设已完成文档嵌入存储) embeddings = HuggingFaceEmbeddings(model_name="all-MiniLM-L6-v2") vector_store = FAISS.load_local("faiss_index", embeddings) retriever = vector_store.as_retriever(search_kwargs={"k": 3}) def process_query(user_query): # 检索相关文档 docs = retriever.get_relevant_documents(user_query) context = "" if docs: context = "相关文档片段:\n" + "\n".join([doc.page_content for doc in docs]) + "\n" # 构建对话链并生成回答 conversation_chain = ConversationChain( llm=pipe, memory=memory, prompt=llama2_prompt.partial(context=context), verbose=False ) return conversation_chain.predict(input=user_query) # 使用示例 print(process_query("文档中关于XX模块的配置说明是什么?")) print(process_query("推荐几个适合周末的户外活动"))
方式二:用RetrievalQA结合记忆(统一处理)
通过RetrievalQA链整合检索与记忆,当无相关文档时 fallback 到普通对话:
from langchain.chains import RetrievalQA from langchain.chains.question_answering import load_qa_chain # 定义QA专用Prompt(适配Llama-2格式) qa_prompt = PromptTemplate( input_variables=["question", "context", "history"], template="""<s>[INST] 你是专业助手,结合对话历史与文档片段回答问题: 对话历史:{history} 文档片段:{context} 问题:{question}[/INST]""" ) # 加载QA链 qa_chain = load_qa_chain(pipe, chain_type="stuff", prompt=qa_prompt) # 构建带记忆的检索QA链 retrieval_qa = RetrievalQA( retriever=retriever, combine_documents_chain=qa_chain, memory=memory, verbose=False, return_source_documents=True ) def handle_query(user_query): result = retrieval_qa({"query": user_query}) # 判断是否有有效检索结果 if not result["source_documents"] or all(len(doc.page_content.strip()) == 0 for doc in result["source_documents"]): # 无相关文档,触发普通对话 conv_chain = ConversationChain(llm=pipe, memory=memory, prompt=llama2_prompt.partial(context=""), verbose=False) return conv_chain.predict(input=user_query) else: return result["result"]
4. 关键注意事项
- 必须设置tokenizer的
pad_token=tokenizer.eos_token,否则会出现生成异常。 - ConversationSummaryBufferMemory的
max_token_limit需根据Llama-2的上下文窗口(13b为4096)合理设置,避免超出模型token限制。 - Prompt中的规则要明确,让模型清晰区分文档依赖回答与普通聊天场景。
内容的提问来源于stack exchange,提问作者AndyLinOuO
相关产品推荐
相关产品推荐

