LangChain RouterChain多输入变量适配问题求助
问题:LangChain RouterChain路由Retrieval Chain时的多输入变量适配错误
问题场景
在LangChain中实现RouterChain,包含Retrieval Chain和标准LLM Chain两个子链:
- 单独运行General Chain(标准LLM Chain)和Specific Chain(Retrieval Chain)均正常
- 通过RouterChain调用General Chain正常,但路由到Retrieval Chain时触发错误
错误日志显示Router传递给Specific Chain的参数为{'input': 'xxx'},但Retrieval Chain需要的输入变量是context和question(或query),参数不匹配导致报错。
核心原因
RouterChain默认将用户输入封装为input变量传递给目标链,但Retrieval Chain(基于RetrievalQA)的标准输入变量是query,且其内部Prompt模板依赖context和question变量,两者的输入参数体系不兼容。
解决方案
步骤1:修正Prompt模板与输入变量对齐
修改Retrieval Chain的Prompt模板,将问题变量改为input(匹配Router传递的参数),同时保留context变量:
template_specific = """Example template Context: {context} Question: {input} """ def set_custom_prompt(): """Prompt template for QA retrieval""" prompt = PromptTemplate(template=template_specific, input_variables=['context', 'input']) return prompt
步骤2:创建参数映射包装链
由于RetrievalQA默认接收query作为输入变量,需要将Router传递的input转换为query,通过TransformChain实现参数映射后,与Retrieval Chain组合为顺序链:
from langchain.chains import TransformChain, SequentialChain # 定义参数转换逻辑:将input转为query def map_input_to_query(inputs: dict) -> dict: return {"query": inputs["input"]} # 创建转换链 transform_chain = TransformChain( input_variables=["input"], output_variables=["query"], transform=map_input_to_query ) # 获取初始化后的Retrieval Chain retrieval_chain = qa_bot() # 组合为顺序链:先转换参数,再调用Retrieval Chain combined_specific_chain = SequentialChain( chains=[transform_chain, retrieval_chain], input_variables=["input"], output_variables=["result"], # RetrievalQA的默认输出键为result verbose=True )
步骤3:更新目标链配置
将包装后的Specific Chain替换原有的未定义llm_chain_specific,加入到destination_chains中:
destination_chains = {} # 初始化General Chain prompt_general = PromptTemplate(template=template_general, input_variables=["input"]) llm_chain_general = LLMChain(prompt=prompt_general, llm=load_llm()) destination_chains["llm_chain_general"] = llm_chain_general # 加入包装后的Specific Chain destination_chains["llm_chain_specific"] = combined_specific_chain
步骤4:对齐自定义RouterChain的输出键
确保自定义MultitypeDestRouteChain的输出键与目标链的输出一致:
class MultitypeDestRouteChain(MultiRouteChain) : """A multi-route chain that uses an LLM router chain to choose amongst prompts.""" router_chain: RouterChain destination_chains: Mapping[str, Chain] default_chain: LLMChain @property def output_keys(self) -> List[str]: return ["result"] # 与combined_specific_chain的输出键对齐
步骤5:修正代码笔误
原代码中set_custom_prompt函数使用了未定义的template2,需替换为template_specific。
修改后的测试调用
# 初始化RouterChain destinations = [f"{p['name']}: {p['description']}" for p in prompt_infos] destinations_str = "\n".join(destinations) router_template = MULTI_PROMPT_ROUTER_TEMPLATE.format(destinations=destinations_str) router_prompt = PromptTemplate( template=router_template, input_variables=['input'], output_parser=RouterOutputParser(), ) router_chain = LLMRouterChain.from_llm(load_llm(), router_prompt) default_chain = ConversationChain(llm=load_llm(), output_key="result") # 初始化最终RouterChain chain = MultitypeDestRouteChain ( router_chain=router_chain, destination_chains=destination_chains, default_chain=default_chain, verbose=True, ) # 测试路由到Retrieval Chain result = chain("This would be a question for example chain?") print(result['result'])
内容的提问来源于stack exchange,提问作者LazyTurtle
相关产品推荐
相关产品推荐

