from airflow import DAG from airflow.operators.python import PythonOperator import pendulum from datetime import timedelta WORLD_SIZE = 2 MASTER_ADDR = "127.0.0.1" MASTER_PORT = "29500" default_args = { "owner": "airflow", "retries": 1, "retry_delay": timedelta(minutes=2), } # ========================= # SINGLE SAFE DDP LAUNCHER # ========================= def run_ddp_job(): import os import torch import torch.distributed as dist import torch.multiprocessing as mp import torch.nn as nn import torch.optim as optim def worker(rank, world_size): os.environ["MASTER_ADDR"] = MASTER_ADDR os.environ["MASTER_PORT"] = MASTER_PORT os.environ["WORLD_SIZE"] = str(world_size) os.environ["RANK"] = str(rank) # optional debug os.environ["NCCL_DEBUG"] = "INFO" os.environ["NCCL_ASYNC_ERROR_HANDLING"] = "1" torch.cuda.set_device(0) dist.init_process_group( backend="nccl", init_method="env://", rank=rank, world_size=world_size, ) model = nn.Linear(10, 10).cuda() ddp_model = torch.nn.parallel.DistributedDataParallel( model, device_ids=[0] ) loss_fn = nn.MSELoss() opt = optim.SGD(ddp_model.parameters(), lr=0.01) for epoch in range(5): x = torch.randn(32, 10).cuda() y = torch.randn(32, 10).cuda() opt.zero_grad() out = ddp_model(x) loss = loss_fn(out, y) loss.backward() opt.step() print(f"rank {rank} epoch {epoch} loss {loss.item()}") dist.destroy_process_group() # IMPORTANT: THIS FIXES NCCL HANG mp.spawn(worker, args=(WORLD_SIZE,), nprocs=WORLD_SIZE, join=True) return {"status": "success"} # ========================= # AIRFLOW DAG # ========================= with DAG( dag_id="pytorch_ddp_airflow_fixed_stable", default_args=default_args, schedule=None, start_date=pendulum.today("UTC").add(days=-1), catchup=False, max_active_runs=1, tags=["ddp", "gpu", "stable"], ) as dag: train = PythonOperator( task_id="train_ddp", python_callable=run_ddp_job, queue="gpu", )