LangChain整合Retriever、Memory与map_reduce链的技术咨询
构建整合Retriever、Memory与map_reduce的链
ConversationalRetrievalChain本身支持同时整合Retriever、Memory和指定chain_type,你之前尝试的ConversationalRetrievalChain.from_llm可以直接实现需求,只需显式指定chain_type='map_reduce'参数即可。
示例代码:
from langchain.chat_models import ChatOpenAI from langchain.vectorstores import Chroma from langchain.memory import ConversationBufferMemory from langchain.chains import ConversationalRetrievalChain # 初始化LLM(替换为你使用的模型) llm = ChatOpenAI(temperature=0) # 初始化Retriever(这里以Chroma向量库为例,替换成你的向量存储实现) vectorstore = Chroma(persist_directory="./chroma_db", embedding_function=your_embedding_model) retriever = vectorstore.as_retriever(search_kwargs={"k": 4}) # 初始化对话记忆 memory = ConversationBufferMemory( memory_key="chat_history", return_messages=True ) # 构建整合链 conversational_chain = ConversationalRetrievalChain.from_llm( llm=llm, retriever=retriever, memory=memory, chain_type="map_reduce", verbose=True # 可选,开启后可查看链执行细节 ) # 测试对话 response = conversational_chain({"question": "LangChain的核心组件有哪些?"}) print(response["answer"])
这个链会自动利用Retriever检索相关文档,通过Memory保留对话上下文,并使用map_reduce模式处理文档:先对单篇文档生成独立结果,再合并所有结果输出最终回答。
仅在令牌数超限时启用map_reduce模式
要实现这个动态切换逻辑,需要先计算检索到的文档总令牌数,再根据阈值选择对应的chain_type。具体步骤如下:
- 实现令牌计数函数,用于统计文档内容的令牌数;
- 检索相关文档并计算总令牌数;
- 根据令牌数是否超过阈值,选择
map_reduce或更高效的stuff模式。
示例代码:
import tiktoken from langchain.chat_models import ChatOpenAI from langchain.vectorstores import Chroma from langchain.memory import ConversationBufferMemory from langchain.chains import ConversationalRetrievalChain # 令牌计数函数(适配gpt-3.5-turbo,可根据你的模型调整) def calculate_doc_tokens(docs): encoder = tiktoken.encoding_for_model("gpt-3.5-turbo") total_tokens = 0 for doc in docs: total_tokens += len(encoder.encode(doc.page_content)) return total_tokens # 设置令牌阈值(根据模型上下文窗口调整,例如gpt-3.5-turbo设为3000) TOKEN_LIMIT = 3000 # 初始化基础组件 llm = ChatOpenAI(temperature=0) vectorstore = Chroma(persist_directory="./chroma_db", embedding_function=your_embedding_model) retriever = vectorstore.as_retriever(search_kwargs={"k": 4}) memory = ConversationBufferMemory(memory_key="chat_history", return_messages=True) # 动态选择chain_type user_question = "请详细解释LangChain中Retriever的工作原理" retrieved_docs = retriever.get_relevant_documents(user_question) total_tokens = calculate_doc_tokens(retrieved_docs) selected_chain_type = "map_reduce" if total_tokens > TOKEN_LIMIT else "stuff" # 构建链 dynamic_chain = ConversationalRetrievalChain.from_llm( llm=llm, retriever=retriever, memory=memory, chain_type=selected_chain_type ) # 执行对话 response = dynamic_chain({"question": user_question}) print(response["answer"])
这种方式既保证了令牌数在限制内时的处理效率,又能在文档内容过多时自动切换到map_reduce模式避免令牌溢出。
内容的提问来源于stack exchange,提问作者tpaus
相关产品推荐
相关产品推荐

