优化LangChain的Map Reduce Documents Chain,缓存Map阶段中间结果
解决LangChain Map Reduce Documents Chain缓存Map阶段结果的方案
核心思路
直接拆分Map Reduce的两个阶段:先单独执行固定的Map模板生成所有文档的处理结果并缓存,之后针对每个不同的Reduce提示,直接复用缓存的Map结果,跳过重复计算。不需要自定义复杂链,用LangChain原生组件即可实现。
具体实现步骤
- 初始化基础组件:加载并分割长文档,定义固定的Map提示模板,初始化LLM。
- 单独执行Map阶段:遍历所有分割后的文档,用固定Map模板生成每个文档的处理结果,将结果存入列表缓存。
- 循环处理不同Reduce提示:针对每个Reduce提示,构建Reduce链,直接传入缓存的Map结果执行,无需重复跑Map。
完整代码示例
from langchain.chat_models import ChatOpenAI from langchain.prompts import PromptTemplate from langchain.chains.summarize import load_summarize_chain from langchain.text_splitter import RecursiveCharacterTextSplitter from langchain.docstore.document import Document # 1. 初始化基础组件 # 加载并分割长文档(示例用模拟文档,替换为实际内容) long_document = """这里替换成你的长文档内容...""" text_splitter = RecursiveCharacterTextSplitter(chunk_size=1000, chunk_overlap=200) docs = text_splitter.split_text(long_document) split_docs = [Document(page_content=text) for text in docs] # 固定的Map提示模板 MAP_PROMPT = PromptTemplate( template="请总结以下文档片段的核心内容:\n\n{text}\n\n总结:", input_variables=["text"] ) llm = ChatOpenAI(temperature=0, model_name="gpt-3.5-turbo") # 2. 执行并缓存Map阶段结果 # 创建临时Map链,仅输出每个文档的Map处理结果 map_chain = load_summarize_chain( llm, chain_type="map_reduce", map_prompt=MAP_PROMPT, reduce_prompt=PromptTemplate(template="{text}", input_variables=["text"]) # 空Reduce模板,直接返回Map结果集合 ) # 执行Map阶段,获取所有文档的输出 raw_map_results = map_chain.run(split_docs) # 根据实际输出格式拆分单个文档的Map结果(示例用换行分割,可根据实际调整) cached_map_results = [res.strip() for res in raw_map_results.split("\n\n") if res.strip()] # 转为Document对象,适配Reduce链的输入要求 cached_map_docs = [Document(page_content=res) for res in cached_map_results] # 3. 循环处理不同Reduce提示 # 示例不同业务场景的Reduce提示 reduce_prompts = [ PromptTemplate(template="请将以下片段整合成一篇简洁的全局摘要:\n\n{text}\n\n最终摘要:", input_variables=["text"]), PromptTemplate(template="从以下片段中提取所有关键数据点并以列表展示:\n\n{text}\n\n关键数据:", input_variables=["text"]), PromptTemplate(template="基于以下片段撰写一篇观点明确的行业评论:\n\n{text}\n\n评论:", input_variables=["text"]) ] for idx, reduce_prompt in enumerate(reduce_prompts): # 针对当前提示构建Reduce链,用stuff模式直接处理缓存结果(数据量大时可换map_reduce/refine模式) reduce_chain = load_summarize_chain( llm, chain_type="stuff", prompt=reduce_prompt ) # 复用缓存的Map结果执行Reduce result = reduce_chain.run(cached_map_docs) print(f"=== 第{idx+1}个Reduce提示结果 ===") print(result) print("\n")
注意事项
- 若Map阶段的输出格式不是换行分割,需根据实际返回内容调整拆分逻辑,确保每个文档的Map结果独立。
- 当缓存的Map结果数量过多时,可将Reduce链的
chain_type改为map_reduce或refine,避免单条提示超出Token限制。 - 可将缓存的Map结果写入本地文件或数据库,实现跨会话的持久化复用。
内容的提问来源于stack exchange,提问作者Frank Rettig
相关产品推荐
相关产品推荐

