diff --git a/dags/test-train-pytorch.py b/dags/test-train-pytorch.py index 94fa594..8f2e16f 100644 --- a/dags/test-train-pytorch.py +++ b/dags/test-train-pytorch.py @@ -2,219 +2,95 @@ 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_ADDR = "127.0.0.1" 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), + "owner": "airflow", + "retries": 1, + "retry_delay": timedelta(minutes=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 +# ========================= +# 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 - from datetime import datetime - print("=" * 60) - print(f"Rank {rank}/{world_size} START {datetime.now()}") - print(f"Hostname: {socket.gethostname()}") - print("=" * 60) + 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) - # ========================= - # STEP 1: SAFE BARRIER - # ========================= - print(f"[{rank}] entering barrier sync...") + # optional debug + os.environ["NCCL_DEBUG"] = "INFO" + os.environ["NCCL_ASYNC_ERROR_HANDLING"] = "1" - with open(BARRIER_FILE, "a+") as f: - f.write(f"{rank}\n") + torch.cuda.set_device(0) - # wait for all ranks - while True: - try: - with open(BARRIER_FILE, "r") as f: - ready = set(f.read().strip().splitlines()) + dist.init_process_group( + backend="nccl", + init_method="env://", + rank=rank, + world_size=world_size, + ) - print(f"[{rank}] ready workers: {ready}") + model = nn.Linear(10, 10).cuda() - if len(ready) == world_size: - print(f"[{rank}] ALL WORKERS READY") - break + ddp_model = torch.nn.parallel.DistributedDataParallel( + model, + device_ids=[0] + ) - except FileNotFoundError: - pass + loss_fn = nn.MSELoss() + opt = optim.SGD(ddp_model.parameters(), lr=0.01) - time.sleep(1) + for epoch in range(5): + x = torch.randn(32, 10).cuda() + y = torch.randn(32, 10).cuda() - 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() + opt.zero_grad() out = ddp_model(x) - loss = criterion(out, y) + loss = loss_fn(out, y) loss.backward() - optimizer.step() + opt.step() - loss_sum += loss.item() + print(f"rank {rank} epoch {epoch} loss {loss.item()}") - print(f"[{rank}] epoch {epoch} loss {loss_sum/5:.4f}") + dist.destroy_process_group() - dist.destroy_process_group() + # IMPORTANT: THIS FIXES NCCL HANG + mp.spawn(worker, args=(WORLD_SIZE,), nprocs=WORLD_SIZE, join=True) - 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"} + return {"status": "success"} # ========================= -# DAG +# AIRFLOW DAG # ========================= with DAG( - dag_id='pytorch_ddp_airflow_fixed', + dag_id="pytorch_ddp_airflow_fixed_stable", default_args=default_args, schedule=None, - start_date=pendulum.today('UTC').add(days=-1), + start_date=pendulum.today("UTC").add(days=-1), catchup=False, max_active_runs=1, - tags=['gpu', 'ddp', 'fixed'], + tags=["ddp", "gpu", "stable"], ) as dag: - prep = PythonOperator( - task_id='prepare_training', - python_callable=prepare_training_func, - queue='gpu', + train = PythonOperator( + task_id="train_ddp", + python_callable=run_ddp_job, + 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