能否将多个SQLDatabaseChain与MultiPromptChain结合使用?及报错排查
问题:SQLDatabaseChain与MultiPromptChain结合时触发ValidationError
尝试将多个SQLDatabaseChain与MultiPromptChain结合使用时,运行代码出现ValidationError,提示destination_chains中的prompt和llm字段为None,不符合验证要求。
原代码示例
table_template = """template 2""" ans_template = """ template 1""" prompt_infos = [ { "name": "table_format", "description": "Good for answering questions if user asks to generate a table", "prompt_template": table_template }, { "name": "ans_format", "description": "Good for answering questions if user don't asks for any specific format", "prompt_template": ans_template } ] llm = OpenAI(temperature=0, model="text-davinci-003", max_tokens=1000) sqlalchemy_url = f'sqlite:///../../../../notebooks/Chinook.db' db = SQLDatabase.from_uri(sqlalchemy_url, view_support=True) destination_chains = {} for p_info in prompt_infos: name = p_info["name"] prompt_template = p_info["prompt_template"] prompt = PromptTemplate(template=prompt_template, input_variables=["input"], validate_template=False) chain = SQLDatabaseChain.from_llm(llm, db, verbose=True, prompt=prompt, use_query_checker=True, top_k=20) destination_chains[name] = chain default_chain = ConversationChain(llm=llm, output_key="text") 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(), validate_template=False ) router_chain = LLMRouterChain.from_llm(llm, router_prompt) chain = MultiPromptChain(router_chain=router_chain, destination_chains=destination_chains, default_chain=default_chain, verbose=True) print(chain.run("Give me top 5 stock codes in a table format"))
报错信息
ValidationError: 10 validation errors for MultiPromptChain destination_chains -> table_format -> prompt none is not an allowed value (type=type_error.none.not_allowed) destination_chains -> table_format -> llm none is not an allowed value (type=type_error.none.not_allowed) ...
问题原因
MultiPromptChain对传入的destination_chains有严格类型要求,它期望所有目标链都是LLMChain类型(或具备LLMChain的核心属性:prompt和llm)。而SQLDatabaseChain是LangChain中的专用链,内部未直接暴露prompt和llm这两个公共属性,导致验证器判定字段为None,触发错误。
解决方案
方案一:包装SQLDatabaseChain适配MultiPromptChain
通过自定义包装类,将SQLDatabaseChain封装成符合MultiPromptChain要求的结构,暴露必要属性:
from langchain.chains import LLMChain # 定义包装类,让SQLDatabaseChain通过MultiPromptChain的验证 class SQLChainWrapper(LLMChain): def __init__(self, sql_chain, llm, prompt, **kwargs): super().__init__(llm=llm, prompt=prompt, **kwargs) self.sql_chain = sql_chain def _call(self, inputs): # 直接复用SQLDatabaseChain的执行逻辑 return self.sql_chain(inputs) # 重构目标链构建逻辑 destination_chains = {} for p_info in prompt_infos: name = p_info["name"] prompt_template = p_info["prompt_template"] prompt = PromptTemplate(template=prompt_template, input_variables=["input"], validate_template=False) sql_chain = SQLDatabaseChain.from_llm(llm, db, verbose=True, prompt=prompt, use_query_checker=True, top_k=20) # 用包装类封装SQLDatabaseChain wrapped_chain = SQLChainWrapper(sql_chain=sql_chain, llm=llm, prompt=prompt, output_key="text") destination_chains[name] = wrapped_chain # 后续RouterChain和MultiPromptChain代码保持不变 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(), validate_template=False ) router_chain = LLMRouterChain.from_llm(llm, router_prompt) chain = MultiPromptChain(router_chain=router_chain, destination_chains=destination_chains, default_chain=default_chain, verbose=True) print(chain.run("Give me top 5 stock codes in a table format"))
方案二:手动实现路由逻辑(更灵活)
放弃使用MultiPromptChain,直接用LLMRouterChain解析路由结果,手动调用对应SQLDatabaseChain,避开类型验证限制:
# 构建路由链 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(), validate_template=False ) router_chain = LLMRouterChain.from_llm(llm, router_prompt) # 自定义路由执行函数 def run_route_query(input_text): # 获取并解析路由结果 route_output = router_chain.predict(input=input_text) parsed_result = RouterOutputParser().parse(route_output) # 根据路由结果调用对应链 if parsed_result.destination in destination_chains: return destination_chains[parsed_result.destination].run(input_text) else: return default_chain.run(input_text) # 测试调用 print(run_route_query("Give me top 5 stock codes in a table format"))
总结
方案二更推荐,无需修改链结构,直接绕过MultiPromptChain的类型限制,逻辑清晰且灵活性更高。
内容的提问来源于stack exchange,提问作者Kamran Khan Khilji
相关产品推荐
相关产品推荐

