Airflow调度DAG参数传递异常:仅首次运行生效,后续调度无参数
问题:Airflow调度运行的DAG无法获取外部触发时传入的参数
外部触发DAG时通过--conf传入参数,首次运行正常,但后续3分钟一次的调度运行均因KeyError: 'lr'失败。查看运行记录发现仅首次外部触发的Run带有配置信息,调度生成的Run的conf为空。
原因
Airflow的调度触发DAG Run不会自动继承外部触发时传入的conf参数,只有手动/外部触发的Run才会携带自定义conf,因此调度运行时context["dag_run"].conf为空,直接通过键取值会触发KeyError。
解决方案
方案1:为参数设置默认值(快速修复)
在获取参数时使用字典的get方法,同时指定默认值,即使conf为空也不会报错:
def train(**context): ti = context["ti"] train = ti.xcom_pull(task_ids="load_data", key="train_path") model_path = ti.xcom_pull(task_ids="load_data", key="model_path") # 使用get方法设置默认值,避免KeyError lr = context["dag_run"].conf.get("lr", 0.001) # 默认学习率0.001 epochs = context["dag_run"].conf.get("epochs", 1) # 默认训练轮次1 name = context["dag_run"].conf.get("name", "default_trial") # 默认任务名称 print(lr) print(epochs) # 补充model_final_name的生成逻辑 model_final_name = f"{model_path}_{name}" ti.xcom_push(key="model_name", value=model_final_name)
方案2:使用Airflow Variables管理参数(灵活可控)
如果需要随时修改参数且所有Run共用,将参数存入Airflow Variables:
- 在Airflow UI的Admin > Variables中添加
lr、epochs、name三个变量 - 修改任务代码读取Variables:
from airflow.models import Variable def train(**context): ti = context["ti"] train = ti.xcom_pull(task_ids="load_data", key="train_path") model_path = ti.xcom_pull(task_ids="load_data", key="model_path") # 从Airflow Variables读取参数 lr = float(Variable.get("lr", default_var=0.001)) epochs = int(Variable.get("epochs", default_var=1)) name = Variable.get("name", default_var="default_trial") print(lr) print(epochs) model_final_name = f"{model_path}_{name}" ti.xcom_push(key="model_name", value=model_final_name)
方案3:在DAG定义中固定参数(适合参数不变的场景)
如果参数不需要动态调整,直接在DAG的default_args中定义,所有Run共用该参数:
with DAG( dag_id="project", default_args={"owner": "airflow", "lr": 0.001, "epochs": 1, "name": "trial_3"}, start_date=datetime(2022, 8, 8), schedule_interval=timedelta(minutes=3), tags=["mnist_4"], catchup=False, ) as dag: # ... 其他代码 ... def train(**context): ti = context["ti"] train = ti.xcom_pull(task_ids="load_data", key="train_path") model_path = ti.xcom_pull(task_ids="load_data", key="model_path") # 从default_args读取固定参数 lr = context["default_args"]["lr"] epochs = context["default_args"]["epochs"] name = context["default_args"]["name"] print(lr) print(epochs) model_final_name = f"{model_path}_{name}" ti.xcom_push(key="model_name", value=model_final_name)
注意点
- 原代码中
train函数的model_final_name未定义,需补充生成逻辑才能正常推送XCom - Airflow 2.x版本中
provide_context=True已被弃用,推荐直接使用**context解构参数
内容的提问来源于stack exchange,提问作者shubh gupta
相关产品推荐
相关产品推荐

