Langchain多输入SequentialChain使用故障求助(附代码)
多输入SequentialChain使用问题排查与正确实现
问题描述
刚接触LangChain,需要实现以下流程:
- 第一个链接收CSV文件片段及该CSV的来源描述,输出可提取的指标列表
- 第二个链接收CSV文件片段和第一个链的输出指标,生成对应的Python脚本
非链式版本代码可正常运行,但自行编写的SequentialChain版本无法执行,需要排查问题并给出正确实现。
代码问题分析
你的SequentialChain代码存在以下几个关键问题:
- 错误使用SimpleSequentialChain:SimpleSequentialChain仅支持单输入→单输出的线性传递,无法处理第二个链需要多个输入(CSV片段+指标)的场景,应使用
SequentialChain替代。 - 链初始化参数错误:
input_variables应指定初始输入的变量名(如data_snippet、source_desc),而非临时变量data_snippet_str- SimpleSequentialChain不支持
output_variables参数,该参数属于SequentialChain的配置项
- run方法调用错误:链的run方法需传入字典格式的参数,而非列表,且无需传入chain对象。
- 未实现来源描述输入:第一个链的prompt未包含CSV来源描述的输入变量,不符合需求。
正确实现代码
from langchain.llms import OpenAI from langchain.prompts import PromptTemplate from langchain.chains import LLMChain, SequentialChain # 初始化LLM llm = OpenAI(model_name="text-davinci-003", openai_api_key='YOUR_API_KEY', temperature=0, max_tokens=3000) # 第一个链:提取指标(包含CSV来源描述输入) prompt_extract = PromptTemplate( input_variables=["data_snippet", "source_desc"], template="""你是一名专注于商业分析的数据科学家,能够从各类数据文件中提取最相关的指标,并能详细说明这些指标如何用于盈利。 数据来源描述:{source_desc} 以下是完整CSV文件的片段作为上下文: {data_snippet} 你的任务是列出可从完整CSV中提取的所有类型的指标,无需进行计算,同时需包含可用于对比的指标。 在指标列表之后,写出列名列表。 可提取的指标及列名: """ ) extract_chain = LLMChain(llm=llm, prompt=prompt_extract, output_key="metrics") # 第二个链:生成Python脚本 prompt_script = PromptTemplate( input_variables=["data_snippet", "metrics"], template="""你是一名专注于商业分析的数据科学家,能够编写强大高效的Python代码从数据集中提取指标。 你的任务是基于以下数据集结构和提取的指标,编写Python脚本: 数据集片段:{data_snippet} 需提取的指标:{metrics} 脚本要求: 1. 使用pandas库 2. 打印所有计算出的指标,以及每个产品及其全部指标 3. 打印完成后,按照示例格式将结果存入metrics_result变量 4. 替换列中的无用字符,计算前转换值为所需类型,注意列名准确性 5. 代码必须以`import pandas as pd`开头,读取CSV文件的逻辑需保留 Python脚本: """ ) script_chain = LLMChain(llm=llm, prompt=prompt_script, output_key="python_script") # 构建SequentialChain:定义输入变量、链顺序、输出变量 overall_chain = SequentialChain( chains=[extract_chain, script_chain], input_variables=["data_snippet", "source_desc"], output_variables=["metrics", "python_script"], verbose=True ) # 示例调用:读取CSV片段并传入参数 def read_csv_data(file_path): import pandas as pd df = pd.read_csv(file_path) # 返回前几行作为片段 return df.head(5).to_string() csv_file_path = "your_csv_file.csv" data_snippet = read_csv_data(csv_file_path) source_desc = "美妆电商产品销售数据CSV,包含产品ID、名称、价格、销量、分类等字段" # 执行链 result = overall_chain.run({ "data_snippet": data_snippet, "source_desc": source_desc }) # 输出结果 print("提取的指标:\n", result["metrics"]) print("\n生成的Python脚本:\n", result["python_script"])
关键说明
- 使用
SequentialChain而非SimpleSequentialChain,支持多输入变量的传递 - 明确定义每个链的
output_key,确保后续链能正确引用前序输出 - 调用链时传入字典格式的初始参数,覆盖所有
input_variables定义的变量 - 完善了第一个链的prompt,加入了CSV来源描述的输入,符合需求
内容的提问来源于stack exchange,提问作者user21412360
相关产品推荐
相关产品推荐

