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

如何在Airflow任务之间传递pandas DataFrame以搭建机器学习pipeline?

Airflow跨任务传递Pandas DataFrame的可行方案

Airflow本身没有限制跨任务传递DataFrame,根据数据体积不同可以选择不同的实现方案:

方案1:小体积DataFrame直接用XCom传递

  • 适用场景:单DataFrame体积小于500MB(可根据你的Airflow元数据库存储上限调整,默认PostgreSQL支持最大1GB的XCom value)
  • 实现逻辑:Airflow 2.x的TaskFlow API默认支持Python对象的pickle序列化,直接在任务中return DataFrame,下游任务直接接收参数即可。注意需要在airflow.cfg中开启enable_xcom_pickling = True(默认关闭,因为pickle存在安全风险,内部可信环境可开启)。
  • 补充说明:生产环境如果存在不可信代码执行的风险,不建议开启pickle,可以改成将DataFrame转成JSON字符串或者base64编码的pickle串之后再存入XCom。

方案2:生产环境通用方案:中间存储+路径传递(最推荐)

  • 适用场景:所有体积的DataFrame,尤其是大体积机器学习数据集
  • 实现逻辑:上游任务把DataFrame写入共享存储(本地共享目录、对象存储、HDFS等),仅把生成的文件路径通过XCom传给下游,下游任务从路径读取DataFrame即可,完全规避XCom的大小限制,性能最优。

拆分任务后的完整代码示例

from datetime import datetime
import pandas as pd
import numpy as np
import os
import lightgbm as lgb
from sklearn.model_selection import StratifiedKFold
from sklearn.metrics import balanced_accuracy_score
from airflow.decorators import dag, task

# 定义共享存储路径,也可以替换为S3路径、HDFS路径等
INTERMEDIATE_STORAGE_PATH = "/tmp/airflow_intermediate"
os.makedirs(INTERMEDIATE_STORAGE_PATH, exist_ok=True)

@dag(dag_id='super_mini_pipeline', schedule_interval=None, 
 start_date=datetime(2021, 11, 5), catchup=False, tags=['ml_pipeline'])
def baseline_pipeline():
    # 任务1:读取原始数据,存为中间parquet返回路径
    @task
    def load_data(label: str):
        path_to_csv = os.path.join('~/airflow/data','leaf.csv') 
        df = pd.read_csv(path_to_csv)
        # 生成唯一的中间文件路径,用时间戳避免不同运行实例冲突
        intermediate_path = os.path.join(INTERMEDIATE_STORAGE_PATH, f"raw_data_{datetime.now().timestamp()}.parquet")
        df.to_parquet(intermediate_path, index=False)
        return intermediate_path

    # 任务2:拆分特征和标签,存为新的中间文件返回路径
    @task
    def split_feature_label(raw_data_path: str, label: str):
        df = pd.read_parquet(raw_data_path)
        y = df[label]
        X = df.drop(label, axis=1)
        X_path = os.path.join(INTERMEDIATE_STORAGE_PATH, f"X_{datetime.now().timestamp()}.parquet")
        y_path = os.path.join(INTERMEDIATE_STORAGE_PATH, f"y_{datetime.now().timestamp()}.parquet")
        X.to_parquet(X_path, index=False)
        y.to_frame().to_parquet(y_path, index=False)
        # 清理上游无用中间文件,可选
        os.remove(raw_data_path)
        return {"X_path": X_path, "y_path": y_path}

    # 任务3:交叉验证训练计算指标
    @task
    def train_and_evaluate(data_paths: dict):
        X = pd.read_parquet(data_paths["X_path"])
        y = pd.read_parquet(data_paths["y_path"]).iloc[:,0]
        folds = StratifiedKFold(n_splits=5, shuffle=True, random_state=10)
        lgbm = lgb.LGBMClassifier(objective='multiclass', random_state=10)
        metrics_lst = []
        for train_idx, val_idx in folds.split(X, y):
            X_train, y_train = X.iloc[train_idx], y.iloc[train_idx]
            X_val, y_val = X.iloc[val_idx], y.iloc[val_idx]
            lgbm.fit(X_train, y_train)
            y_pred = lgbm.predict(X_val)
            cv_balanced_accuracy = balanced_accuracy_score(y_val, y_pred)
            metrics_lst.append(cv_balanced_accuracy)
        avg_performance = np.mean(metrics_lst)
        print(f"Avg Performance: {avg_performance}")
        # 清理中间文件,可选
        os.remove(data_paths["X_path"])
        os.remove(data_paths["y_path"])
        return avg_performance

    # 任务依赖编排
    raw_path = load_data(label='species')
    data_paths = split_feature_label(raw_path, label='species')
    train_and_evaluate(data_paths)

# dag invocation
pipeline_dag = baseline_pipeline()

额外说明

如果你的Airflow是分布式部署,需要保证所有工作节点都能访问到你定义的中间存储路径,用对象存储、NFS共享目录、HDFS都可以实现。小数据场景下你也可以直接把DataFrame转成df.to_dict("records")之后return,不需要存中间文件,Airflow会自动把这个列表序列化存到XCom,下游直接接收后转成DataFrame即可。

内容的提问来源于stack exchange,提问作者Cherry Wu

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.27 06:15:08