from airflow import DAG from airflow.operators.python import PythonOperator import pendulum from datetime import timedelta import os import time # --- CONFIGURATION --- WORLD_SIZE = 2 MASTER_ADDR = "airflow-worker-gpu-0.airflow-worker-gpu" MASTER_PORT = "29500" NCCL_TIMEOUT = 1800 BARRIER_FILE = f"/tmp/ddp_barrier_{WORLD_SIZE}" default_args = { 'owner': 'airflow', 'retries': 1, 'retry_delay': timedelta(minutes=2), 'execution_timeout': timedelta(hours=2), } def prepare_training_func(): # reset barrier try: if os.path.exists(BARRIER_FILE): os.remove(BARRIER_FILE) except: pass print("Barrier initialized") return True def run_training_node_func(rank, world_size): import socket import torch import torch.distributed as dist import torch.nn as nn import torch.optim as optim from datetime import datetime print("=" * 60) print(f"Rank {rank}/{world_size} START {datetime.now()}") print(f"Hostname: {socket.gethostname()}") print("=" * 60) # ========================= # STEP 1: SAFE BARRIER # ========================= print(f"[{rank}] entering barrier sync...") with open(BARRIER_FILE, "a+") as f: f.write(f"{rank}\n") # wait for all ranks while True: try: with open(BARRIER_FILE, "r") as f: ready = set(f.read().strip().splitlines()) print(f"[{rank}] ready workers: {ready}") if len(ready) == world_size: print(f"[{rank}] ALL WORKERS READY") break except FileNotFoundError: pass time.sleep(1) time.sleep(2) # small stabilization delay # ========================= # STEP 2: NCCL ENV # ========================= os.environ['MASTER_ADDR'] = MASTER_ADDR os.environ['MASTER_PORT'] = MASTER_PORT os.environ['WORLD_SIZE'] = str(world_size) os.environ['RANK'] = str(rank) # IMPORTANT FIX (IB stability toggle) os.environ["NCCL_IB_DISABLE"] = "1" # <- FIX #2 (can turn OFF later) os.environ["NCCL_SOCKET_IFNAME"] = "eth0" os.environ['NCCL_DEBUG'] = 'INFO' os.environ['NCCL_TIMEOUT'] = str(NCCL_TIMEOUT) os.environ['NCCL_BLOCKING_WAIT'] = '1' print(f"[{rank}] env ready") # ========================= # STEP 3: INIT PROCESS GROUP # ========================= dist.init_process_group( backend="nccl", init_method="env://", timeout=timedelta(minutes=10), rank=rank, world_size=world_size ) print(f"[{rank}] process group OK") # ========================= # STEP 4: GPU # ========================= device = torch.device("cuda:0") torch.cuda.set_device(device) model = nn.Sequential( nn.Linear(10, 128), nn.ReLU(), nn.Linear(128, 10) ).to(device) ddp_model = nn.parallel.DistributedDataParallel( model, device_ids=[0], output_device=0 ) criterion = nn.MSELoss() optimizer = optim.SGD(ddp_model.parameters(), lr=0.001) # ========================= # STEP 5: TRAIN LOOP # ========================= for epoch in range(5): ddp_model.train() loss_sum = 0 for _ in range(5): x = torch.randn(32, 10).to(device) y = torch.randn(32, 10).to(device) optimizer.zero_grad() out = ddp_model(x) loss = criterion(out, y) loss.backward() optimizer.step() loss_sum += loss.item() print(f"[{rank}] epoch {epoch} loss {loss_sum/5:.4f}") dist.destroy_process_group() print(f"[{rank}] DONE") return {"rank": rank, "status": "success"} def cleanup_sync_state_func(): try: if os.path.exists(BARRIER_FILE): os.remove(BARRIER_FILE) except: pass print("Barrier cleaned") return True def training_summary_func(**context): ti = context['ti'] print("\n=== SUMMARY ===") for rank in range(WORLD_SIZE): res = ti.xcom_pull(task_ids=f"train_rank_{rank}") print(f"Rank {rank}: {res}") return {"status": "done"} # ========================= # DAG # ========================= with DAG( dag_id='pytorch_ddp_airflow_fixed', default_args=default_args, schedule=None, start_date=pendulum.today('UTC').add(days=-1), catchup=False, max_active_runs=1, tags=['gpu', 'ddp', 'fixed'], ) as dag: prep = PythonOperator( task_id='prepare_training', python_callable=prepare_training_func, queue='gpu', ) training_tasks = [] for i in range(WORLD_SIZE): t = PythonOperator( task_id=f'train_rank_{i}', python_callable=run_training_node_func, op_kwargs={'rank': i, 'world_size': WORLD_SIZE}, queue='gpu', ) training_tasks.append(t) cleanup = PythonOperator( task_id='cleanup', python_callable=cleanup_sync_state_func, queue='gpu', trigger_rule='all_done', ) summary = PythonOperator( task_id='summary', python_callable=training_summary_func, trigger_rule='all_done', ) prep >> training_tasks >> cleanup >> summary