You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

能否将多个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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.18 15:54:55