diff --git a/dags/test-train-pytorch.py b/dags/test-train-pytorch.py index 1263714..939d95e 100644 --- a/dags/test-train-pytorch.py +++ b/dags/test-train-pytorch.py @@ -1,115 +1,273 @@ from airflow import DAG -from airflow.providers.standard.operators.python import PythonOperator +from airflow.operators.python import PythonOperator +from airflow.models import Variable import pendulum from datetime import timedelta +import json +# --- CONFIGURATION --- +WORLD_SIZE = 2 +MASTER_ADDR = "airflow-worker-gpu-0.airflow-worker-gpu" MASTER_PORT = "29500" +NCCL_TIMEOUT = 1800 default_args = { - "owner": "airflow", - "retries": 0, - "execution_timeout": timedelta(minutes=20), + 'owner': 'airflow', + 'retries': 1, + 'retry_delay': timedelta(minutes=2), + 'execution_timeout': timedelta(hours=2), } -# ------------------------- -# RANK 0 -# ------------------------- -def rank0(**context): - import os, socket, torch - import torch.distributed as dist - - addr = socket.gethostname() + ".airflow-worker-gpu" - - context["ti"].xcom_push(key="master_addr", value=addr) - - os.environ.update({ - "MASTER_ADDR": addr, - "MASTER_PORT": MASTER_PORT, - "WORLD_SIZE": "2", - "RANK": "0", - - "NCCL_DEBUG": "INFO", - "NCCL_SOCKET_IFNAME": "eth0", - - # IB - "NCCL_IB_DISABLE": "0", - "NCCL_IB_GID_INDEX": "0", - }) - - dist.init_process_group("nccl") - - torch.cuda.set_device(0) - - x = torch.ones(1, device="cuda") * 1 - dist.all_reduce(x) - - print(f"[rank0] result = {x.item()}") - - dist.destroy_process_group() +def prepare_training_func(): + """Initialize shared state for synchronization""" + sync_state = { + 'ready_workers': [], + 'training_started': False, + 'start_time': None + } + Variable.set('ddp_sync_state', json.dumps(sync_state)) + print("Training preparation complete. Sync state initialized.") + return True -# ------------------------- -# RANK 1 -# ------------------------- -def rank1(**context): - import os, time, torch - import torch.distributed as dist +def run_training_node_func(rank, world_size): + """Execute distributed training with synchronization""" + import os + import socket + import torch + import torch.distributed as dist + import torch.nn as nn + import torch.optim as optim + from datetime import datetime + import time - addr = None - for _ in range(30): - addr = context["ti"].xcom_pull(task_ids="rank0", key="master_addr") - if addr: - break - time.sleep(2) + print(f"{'='*60}") + print(f"Node Rank {rank}/{world_size} - Starting at {datetime.now()}") + print(f"Hostname: {socket.gethostname()}") + print(f"{'='*60}") - os.environ.update({ - "MASTER_ADDR": addr, - "MASTER_PORT": MASTER_PORT, - "WORLD_SIZE": "2", - "RANK": "1", + # STEP 1: Signal that this worker is ready + print(f"[{rank}] Signaling ready state...") + max_wait = 300 # 5 minutes + start_wait = time.time() - "NCCL_DEBUG": "INFO", - "NCCL_SOCKET_IFNAME": "eth0", + while time.time() - start_wait < max_wait: + try: + sync_state = json.loads(Variable.get('ddp_sync_state', default_var='{}')) - "NCCL_IB_DISABLE": "0", - "NCCL_IB_GID_INDEX": "0", - }) + if rank not in sync_state.get('ready_workers', []): + sync_state.setdefault('ready_workers', []).append(rank) + Variable.set('ddp_sync_state', json.dumps(sync_state)) + print(f"[{rank}] Marked as ready. Ready workers: {sync_state['ready_workers']}") - dist.init_process_group("nccl") + # STEP 2: Wait for all workers to be ready + if len(sync_state.get('ready_workers', [])) == world_size: + print(f"[{rank}] All {world_size} workers are ready! Proceeding to training...") + break - torch.cuda.set_device(0) + print(f"[{rank}] Waiting for other workers... ({len(sync_state.get('ready_workers', []))}/{world_size} ready)") + time.sleep(2) - x = torch.ones(1, device="cuda") * 2 - dist.all_reduce(x) + except Exception as e: + print(f"[{rank}] Error during sync: {e}") + time.sleep(2) + else: + raise RuntimeError(f"[{rank}] Timeout waiting for all workers to be ready!") - print(f"[rank1] result = {x.item()}") + # Small delay to ensure all workers see the ready state + time.sleep(3) - dist.destroy_process_group() + # STEP 3: Configure distributed environment + os.environ['MASTER_ADDR'] = MASTER_ADDR + os.environ['MASTER_PORT'] = MASTER_PORT + os.environ['WORLD_SIZE'] = str(world_size) + os.environ['RANK'] = str(rank) + 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}] Environment configured:") + print(f" MASTER_ADDR: {MASTER_ADDR}") + print(f" MASTER_PORT: {MASTER_PORT}") + print(f" RANK: {rank}") + print(f" WORLD_SIZE: {world_size}") + + # STEP 4: Initialize process group + print(f"[{rank}] Initializing process group (backend=nccl)...") + try: + dist.init_process_group( + backend="nccl", + init_method="env://", + timeout=timedelta(minutes=10), + rank=rank, + world_size=world_size + ) + print(f"[{rank}] āœ“ Successfully joined distributed group!") + print(f" Process group size: {dist.get_world_size()}") + print(f" My rank: {dist.get_rank()}") + except Exception as e: + print(f"[{rank}] āœ— Failed to initialize process group!") + print(f" Error: {str(e)}") + raise + + # STEP 5: Setup GPU + device = torch.device("cuda:0") + torch.cuda.set_device(device) + + print(f"[{rank}] GPU Configuration:") + print(f" Device: {torch.cuda.get_device_name(0)}") + print(f" Memory: {torch.cuda.get_device_properties(0).total_memory / 1e9:.2f} GB") + + # STEP 6: Define model + model = nn.Sequential( + nn.Linear(10, 128), + nn.ReLU(), + nn.Linear(128, 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, momentum=0.9) + + print(f"[{rank}] Model initialized with {sum(p.numel() for p in model.parameters())} parameters") + + # STEP 7: Training loop + print(f"\n[{rank}] {'='*60}") + print(f"[{rank}] Starting Training Loop") + print(f"[{rank}] {'='*60}") + + num_epochs = 10 + batch_size = 32 + num_batches = 5 + + for epoch in range(num_epochs): + ddp_model.train() + epoch_loss = 0.0 + + for batch_idx in range(num_batches): + torch.manual_seed(epoch * num_batches + batch_idx) + inputs = torch.randn(batch_size, 10).to(device) + labels = torch.randn(batch_size, 10).to(device) + + optimizer.zero_grad() + outputs = ddp_model(inputs) + loss = criterion(outputs, labels) + loss.backward() + optimizer.step() + + epoch_loss += loss.item() + + avg_loss = epoch_loss / num_batches + + # Synchronize loss across ranks + loss_tensor = torch.tensor([avg_loss]).to(device) + dist.all_reduce(loss_tensor, op=dist.ReduceOp.AVG) + global_avg_loss = loss_tensor.item() + + if rank == 0: + print(f"[{rank}] Epoch {epoch+1}/{num_epochs} | Global Avg Loss: {global_avg_loss:.6f}") + + print(f"\n[{rank}] {'='*60}") + print(f"[{rank}] Training Complete!") + print(f"[{rank}] {'='*60}") + + # STEP 8: Cleanup + dist.destroy_process_group() + print(f"[{rank}] Process group destroyed. Finished at {datetime.now()}") + + return { + 'rank': rank, + 'final_loss': global_avg_loss if rank == 0 else avg_loss, + 'epochs_completed': num_epochs, + 'status': 'success' + } + + +def cleanup_sync_state_func(): + """Clean up synchronization state""" + try: + Variable.delete('ddp_sync_state') + print("Sync state cleaned up") + except: + pass + return True + + +def training_summary_func(**context): + """Aggregate and display training results""" + ti = context['ti'] + + print(f"\n{'='*60}") + print(f"DISTRIBUTED TRAINING SUMMARY") + print(f"{'='*60}") + + for rank in range(WORLD_SIZE): + try: + result = ti.xcom_pull(task_ids=f"train_rank_{rank}") + if result and result.get('status') == 'success': + print(f"Rank {result['rank']}: āœ“ Completed {result['epochs_completed']} epochs") + print(f" Final loss: {result['final_loss']:.6f}") + except Exception as e: + print(f"Warning: Could not get result for rank {rank}: {e}") + + print(f"\nāœ“ Training completed!") + return {'status': 'success', 'workers': WORLD_SIZE} -# ------------------------- -# DAG -# ------------------------- with DAG( - dag_id="ib_simple_test", - start_date=pendulum.today("UTC").add(days=-1), - schedule=None, - catchup=False, - max_active_runs=1, - default_args=default_args, + dag_id='pytorch_distributed_training_ddp_production', + default_args=default_args, + schedule=None, + start_date=pendulum.today('UTC').add(days=-1), + catchup=False, + tags=['gpu', 'ml', 'distributed'], + max_active_runs=1, + max_active_tasks=10, ) as dag: - r0 = PythonOperator( - task_id="rank0", - python_callable=rank0, - queue="gpu", - ) + # Preparation task + prep = PythonOperator( + task_id='prepare_training', + python_callable=prepare_training_func, + queue='gpu' + ) - r1 = PythonOperator( - task_id="rank1", - python_callable=rank1, - queue="gpu", - ) + # Training tasks - one per rank + training_tasks = [] + for i in range(WORLD_SIZE): + task = 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(task) + + # Cleanup task + cleanup = PythonOperator( + task_id='cleanup_sync_state', + python_callable=cleanup_sync_state_func, + queue='gpu', + trigger_rule='all_done' + ) + + # Summary task + summary = PythonOperator( + task_id='training_summary', + python_callable=training_summary_func, + queue='gpu', + trigger_rule='all_done' + ) + + # Set dependencies + prep >> training_tasks >> cleanup >> summary - r0 >> r1