From 1a79410ad7c6303d2eee30f579c721d3d015a462 Mon Sep 17 00:00:00 2001 From: George Stykalin Date: Wed, 15 Apr 2026 00:07:51 +0300 Subject: [PATCH] update dags --- dags/test-train-pytorch.py | 378 +++++++++++++------------------------ 1 file changed, 130 insertions(+), 248 deletions(-) diff --git a/dags/test-train-pytorch.py b/dags/test-train-pytorch.py index 50a4c22..003e2d0 100644 --- a/dags/test-train-pytorch.py +++ b/dags/test-train-pytorch.py @@ -1,295 +1,177 @@ from airflow import DAG 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': 1, - 'retry_delay': timedelta(minutes=2), - 'execution_timeout': timedelta(hours=2), + "owner": "airflow", + "retries": 0, + "execution_timeout": timedelta(hours=2), } +# ------------------------- +# SIMPLE PREP (NO VARIABLES) +# ------------------------- 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 + print("Starting distributed training (no sync state needed)") + return True +# ------------------------- +# CORE TRAINING +# ------------------------- 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 + import os + import time + import torch + import torch.distributed as dist + import torch.nn as nn + import torch.optim as optim + from datetime import datetime + import socket - print(f"{'='*60}") - print(f"Node Rank {rank}/{world_size} - Starting at {datetime.now()}") - print(f"Hostname: {socket.gethostname()}") - print(f"{'='*60}") + print("=" * 80) + print(f"RANK {rank}/{world_size} START {datetime.now()}") + print(f"HOST: {socket.gethostname()}") + print("=" * 80) - # STEP 1: Signal that this worker is ready - print(f"[{rank}] Signaling ready state...") - max_wait = 300 # 5 minutes - start_wait = time.time() + # ------------------------- + # FIXED BARRIER (important) + # ------------------------- + print(f"[{rank}] sync barrier (static sleep)") + time.sleep(10) - while time.time() - start_wait < max_wait: - try: - sync_state = json.loads(Variable.get('ddp_sync_state', default_var='{}')) + # ------------------------- + # STATIC CONFIG (NO AIRFLOW VARIABLES) + # ------------------------- + os.environ["MASTER_ADDR"] = MASTER_ADDR + os.environ["MASTER_PORT"] = MASTER_PORT + os.environ["WORLD_SIZE"] = str(world_size) + os.environ["RANK"] = str(rank) - 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']}") + os.environ["NCCL_SOCKET_IFNAME"] = "eth0" + os.environ["NCCL_DEBUG"] = "INFO" + os.environ["TORCH_NCCL_BLOCKING_WAIT"] = "1" - # 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}] MASTER = {MASTER_ADDR}:{MASTER_PORT}") - print(f"[{rank}] Waiting for other workers... ({len(sync_state.get('ready_workers', []))}/{world_size} ready)") - time.sleep(2) + # ------------------------- + # INIT PROCESS GROUP + # ------------------------- + dist.init_process_group( + backend="nccl", + init_method="env://", + rank=rank, + world_size=world_size, + ) - 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"[{rank}] DDP INIT OK") - # Small delay to ensure all workers see the ready state - time.sleep(3) + # ------------------------- + # GPU BINDING (CRITICAL FIX) + # ------------------------- + torch.cuda.set_device(0) + device = torch.device("cuda:0") - # STEP 3: Configure distributed Environment + print(f"[{rank}] GPU = {torch.cuda.get_device_name(0)}") - import socket + # ------------------------- + # MODEL + # ------------------------- + model = nn.Sequential( + nn.Linear(10, 128), + nn.ReLU(), + nn.Linear(128, 10), + ).to(device) - if rank == 0: - # Rank 0 becomes master - master_addr = socket.gethostbyname(socket.gethostname()) - Variable.set("MASTER_ADDR_DYNAMIC", master_addr) - print(f"[{rank}] Acting as MASTER at {master_addr}") - else: - # Other ranks wait for master - print(f"[{rank}] Waiting for MASTER_ADDR...") - while True: - try: - master_addr = Variable.get("MASTER_ADDR_DYNAMIC", default_var=None) - if master_addr: - break - except: - pass - time.sleep(1) + ddp_model = nn.parallel.DistributedDataParallel( + model, + device_ids=[0], + ) - print(f"[{rank}] Found MASTER at {master_addr}") + optimizer = optim.SGD(ddp_model.parameters(), lr=0.001) + loss_fn = nn.MSELoss() - 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_PORT_RANGE'] = '30000-30100' - os.environ['NCCL_DEBUG'] = 'INFO' - os.environ['NCCL_TIMEOUT'] = str(NCCL_TIMEOUT) - os.environ['NCCL_BLOCKING_WAIT'] = '1' + # ------------------------- + # TRAIN LOOP + # ------------------------- + for epoch in range(5): + x = torch.randn(32, 10).to(device) + y = torch.randn(32, 10).to(device) - 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}") + optimizer.zero_grad() + out = ddp_model(x) + loss = loss_fn(out, y) + loss.backward() + optimizer.step() - # 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 + print(f"[{rank}] epoch={epoch} loss={loss.item():.4f}") - # STEP 5: Setup GPU - device = torch.device("cuda:0") - torch.cuda.set_device(device) + dist.destroy_process_group() - 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' - } + return {"rank": rank, "status": "ok"} -def cleanup_sync_state_func(): - """Clean up synchronization state""" - try: - Variable.delete('ddp_sync_state') - print("Sync state cleaned up") - except: - pass - return True +# ------------------------- +# CLEANUP +# ------------------------- +def cleanup_func(): + print("cleanup done") + 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} +def summary_func(**context): + print("training done") + return True +# ------------------------- +# DAG +# ------------------------- with DAG( - 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, + dag_id="ddp_airflow_stable", + start_date=pendulum.today("UTC").add(days=-1), + schedule=None, + catchup=False, + max_active_runs=1, + max_active_tasks=2, + default_args=default_args, ) as dag: - # Preparation task - prep = PythonOperator( - task_id='prepare_training', - python_callable=prepare_training_func, - queue='gpu' - ) + prep = PythonOperator( + task_id="prep", + python_callable=prepare_training_func, + 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) + tasks = [] + for r in range(WORLD_SIZE): + t = PythonOperator( + task_id=f"train_rank_{r}", + python_callable=run_training_node_func, + op_kwargs={"rank": r, "world_size": WORLD_SIZE}, + queue="gpu", + ) + tasks.append(t) - # Cleanup task - cleanup = PythonOperator( - task_id='cleanup_sync_state', - python_callable=cleanup_sync_state_func, - queue='gpu', - trigger_rule='all_done' - ) + cleanup = PythonOperator( + task_id="cleanup", + python_callable=cleanup_func, + trigger_rule="all_done", + queue="gpu", + ) - # Summary task - summary = PythonOperator( - task_id='training_summary', - python_callable=training_summary_func, - queue='gpu', - trigger_rule='all_done' - ) + summary = PythonOperator( + task_id="summary", + python_callable=summary_func, + trigger_rule="all_done", + queue="gpu", + ) - # Set dependencies - prep >> training_tasks >> cleanup >> summary + prep >> tasks >> cleanup >> summary