From f84c133b002912b8025f069c9205bbff94a4d0e0 Mon Sep 17 00:00:00 2001 From: George Stykalin Date: Wed, 15 Apr 2026 12:12:11 +0300 Subject: [PATCH] update dags --- dags/test-train-pytorch.py | 316 +++++++++++++++++++++++++++++-------- 1 file changed, 254 insertions(+), 62 deletions(-) diff --git a/dags/test-train-pytorch.py b/dags/test-train-pytorch.py index c26197e..720a87f 100644 --- a/dags/test-train-pytorch.py +++ b/dags/test-train-pytorch.py @@ -1,80 +1,272 @@ from airflow import DAG -from airflow.providers.standard.operators.python import PythonOperator -from airflow.operators.bash import BashOperator -from datetime import datetime -import os +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", + 'owner': 'airflow', + 'retries': 1, + 'retry_delay': timedelta(minutes=2), + 'execution_timeout': timedelta(hours=2), } -# ----------------------- -# PREP -# ----------------------- -def prepare_training(): - print("Preparing DDP training environment") - - os.environ["TOKENIZERS_PARALLELISM"] = "false" - os.environ["NCCL_DEBUG"] = "INFO" - - return "prepared" +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 -# ----------------------- -# CLEANUP -# ----------------------- -def cleanup(): - print("Cleaning up after training") - return "cleaned" +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 + + print(f"{'='*60}") + print(f"Node Rank {rank}/{world_size} - Starting at {datetime.now()}") + print(f"Hostname: {socket.gethostname()}") + print(f"{'='*60}") + + # STEP 1: Signal that this worker is ready + print(f"[{rank}] Signaling ready state...") + max_wait = 300 # 5 minutes + start_wait = time.time() + + while time.time() - start_wait < max_wait: + try: + sync_state = json.loads(Variable.get('ddp_sync_state', default_var='{}')) + + 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']}") + + # 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 + + print(f"[{rank}] Waiting for other workers... ({len(sync_state.get('ready_workers', []))}/{world_size} ready)") + time.sleep(2) + + 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!") + + # Small delay to ensure all workers see the ready state + time.sleep(3) + + # 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="pytorch_ddp_airflow_fixed_production", - start_date=datetime(2024, 1, 1), - schedule=None, - catchup=False, - default_args=default_args, - tags=["ddp", "pytorch", "gpu"], + 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: - prepare_training_task = PythonOperator( - task_id="prepare_training", - python_callable=prepare_training, - queue="gpu", # ✅ FIX #1 - ) + # Preparation task + prep = PythonOperator( + task_id='prepare_training', + python_callable=prepare_training_func, + queue='gpu' + ) - train_ddp_task = BashOperator( - task_id="train_ddp", - bash_command=""" - set -e + # 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) - export NCCL_DEBUG=INFO - export NCCL_ASYNC_ERROR_HANDLING=1 - export NCCL_IB_DISABLE=1 + # Cleanup task + cleanup = PythonOperator( + task_id='cleanup_sync_state', + python_callable=cleanup_sync_state_func, + queue='gpu', + trigger_rule='all_done' + ) - echo "Starting torchrun DDP training" + # Summary task + summary = PythonOperator( + task_id='training_summary', + python_callable=training_summary_func, + queue='gpu', + trigger_rule='all_done' + ) - torchrun \ - --nproc_per_node=2 \ - --master_port=29500 \ - /opt/airflow/dags/repo/train.py - - echo "Training finished" - """, - queue="gpu", # ✅ FIX #1 - ) - - cleanup_task = PythonOperator( - task_id="cleanup", - python_callable=cleanup, - queue="gpu", # ✅ FIX #1 - ) - - - # DAG FLOW - prepare_training_task >> train_ddp_task >> cleanup_task + # Set dependencies + prep >> training_tasks >> cleanup >> summary