update dags

This commit is contained in:
2026-04-15 12:12:11 +03:00
parent 6dddc69d11
commit f84c133b00

View File

@@ -1,80 +1,272 @@
from airflow import DAG from airflow import DAG
from airflow.providers.standard.operators.python import PythonOperator from airflow.operators.python import PythonOperator
from airflow.operators.bash import BashOperator from airflow.models import Variable
from datetime import datetime import pendulum
import os 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 = { default_args = {
"owner": "airflow", 'owner': 'airflow',
'retries': 1,
'retry_delay': timedelta(minutes=2),
'execution_timeout': timedelta(hours=2),
} }
# ----------------------- def prepare_training_func():
# PREP """Initialize shared state for synchronization"""
# ----------------------- sync_state = {
def prepare_training(): 'ready_workers': [],
print("Preparing DDP training environment") 'training_started': False,
'start_time': None
os.environ["TOKENIZERS_PARALLELISM"] = "false" }
os.environ["NCCL_DEBUG"] = "INFO" Variable.set('ddp_sync_state', json.dumps(sync_state))
print("Training preparation complete. Sync state initialized.")
return "prepared" return True
# ----------------------- def run_training_node_func(rank, world_size):
# CLEANUP """Execute distributed training with synchronization"""
# ----------------------- import os
def cleanup(): import socket
print("Cleaning up after training") import torch
return "cleaned" 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( with DAG(
dag_id="pytorch_ddp_airflow_fixed_production", dag_id='pytorch_distributed_training_ddp_production',
start_date=datetime(2024, 1, 1), default_args=default_args,
schedule=None, schedule=None,
catchup=False, start_date=pendulum.today('UTC').add(days=-1),
default_args=default_args, catchup=False,
tags=["ddp", "pytorch", "gpu"], tags=['gpu', 'ml', 'distributed'],
max_active_runs=1,
max_active_tasks=10,
) as dag: ) as dag:
prepare_training_task = PythonOperator( # Preparation task
task_id="prepare_training", prep = PythonOperator(
python_callable=prepare_training, task_id='prepare_training',
queue="gpu", # ✅ FIX #1 python_callable=prepare_training_func,
) queue='gpu'
)
train_ddp_task = BashOperator( # Training tasks - one per rank
task_id="train_ddp", training_tasks = []
bash_command=""" for i in range(WORLD_SIZE):
set -e 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 # Cleanup task
export NCCL_ASYNC_ERROR_HANDLING=1 cleanup = PythonOperator(
export NCCL_IB_DISABLE=1 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 \ # Set dependencies
--nproc_per_node=2 \ prep >> training_tasks >> cleanup >> summary
--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