update dags

This commit is contained in:
2026-04-15 11:54:13 +03:00
parent 54ead9f90a
commit 5f7db605a6

View File

@@ -1,9 +1,9 @@
from airflow import DAG from airflow import DAG
from airflow.operators.python import PythonOperator from airflow.operators.python import PythonOperator
from airflow.models import Variable
import pendulum import pendulum
from datetime import timedelta from datetime import timedelta
import json import os
import time
# --- CONFIGURATION --- # --- CONFIGURATION ---
WORLD_SIZE = 2 WORLD_SIZE = 2
@@ -11,263 +11,210 @@ MASTER_ADDR = "airflow-worker-gpu-0.airflow-worker-gpu"
MASTER_PORT = "29500" MASTER_PORT = "29500"
NCCL_TIMEOUT = 1800 NCCL_TIMEOUT = 1800
BARRIER_FILE = f"/tmp/ddp_barrier_{WORLD_SIZE}"
default_args = { default_args = {
'owner': 'airflow', 'owner': 'airflow',
'retries': 1, 'retries': 1,
'retry_delay': timedelta(minutes=2), 'retry_delay': timedelta(minutes=2),
'execution_timeout': timedelta(hours=2), 'execution_timeout': timedelta(hours=2),
} }
def prepare_training_func(): def prepare_training_func():
"""Initialize shared state for synchronization""" # reset barrier
sync_state = { try:
'ready_workers': [], if os.path.exists(BARRIER_FILE):
'training_started': False, os.remove(BARRIER_FILE)
'start_time': None except:
} pass
Variable.set('ddp_sync_state', json.dumps(sync_state))
print("Training preparation complete. Sync state initialized.") print("Barrier initialized")
return True return True
def run_training_node_func(rank, world_size): def run_training_node_func(rank, world_size):
"""Execute distributed training with synchronization""" import socket
import os import torch
import socket import torch.distributed as dist
import torch import torch.nn as nn
import torch.distributed as dist import torch.optim as optim
import torch.nn as nn from datetime import datetime
import torch.optim as optim
from datetime import datetime
import time
print(f"{'='*60}") print("=" * 60)
print(f"Node Rank {rank}/{world_size} - Starting at {datetime.now()}") print(f"Rank {rank}/{world_size} START {datetime.now()}")
print(f"Hostname: {socket.gethostname()}") print(f"Hostname: {socket.gethostname()}")
print(f"{'='*60}") print("=" * 60)
# STEP 1: Signal that this worker is ready # =========================
print(f"[{rank}] Signaling ready state...") # STEP 1: SAFE BARRIER
max_wait = 300 # 5 minutes # =========================
start_wait = time.time() print(f"[{rank}] entering barrier sync...")
while time.time() - start_wait < max_wait: with open(BARRIER_FILE, "a+") as f:
try: f.write(f"{rank}\n")
sync_state = json.loads(Variable.get('ddp_sync_state', default_var='{}'))
if rank not in sync_state.get('ready_workers', []): # wait for all ranks
sync_state.setdefault('ready_workers', []).append(rank) while True:
Variable.set('ddp_sync_state', json.dumps(sync_state)) try:
print(f"[{rank}] Marked as ready. Ready workers: {sync_state['ready_workers']}") with open(BARRIER_FILE, "r") as f:
ready = set(f.read().strip().splitlines())
# STEP 2: Wait for all workers to be ready print(f"[{rank}] ready workers: {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)") if len(ready) == world_size:
time.sleep(2) print(f"[{rank}] ALL WORKERS READY")
break
except Exception as e: except FileNotFoundError:
print(f"[{rank}] Error during sync: {e}") pass
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(1)
time.sleep(3)
# STEP 3: Configure distributed environment time.sleep(2) # small stabilization delay
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}") # STEP 2: NCCL ENV
print(f" MASTER_PORT: {MASTER_PORT}") # =========================
print(f" RANK: {rank}") os.environ['MASTER_ADDR'] = MASTER_ADDR
print(f" WORLD_SIZE: {world_size}") os.environ['MASTER_PORT'] = MASTER_PORT
os.environ['WORLD_SIZE'] = str(world_size)
os.environ['RANK'] = str(rank)
# STEP 4: Initialize process group # IMPORTANT FIX (IB stability toggle)
print(f"[{rank}] Initializing process group (backend=nccl)...") os.environ["NCCL_IB_DISABLE"] = "1" # <- FIX #2 (can turn OFF later)
try: os.environ["NCCL_SOCKET_IFNAME"] = "eth0"
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 os.environ['NCCL_DEBUG'] = 'INFO'
device = torch.device("cuda:0") os.environ['NCCL_TIMEOUT'] = str(NCCL_TIMEOUT)
torch.cuda.set_device(device) os.environ['NCCL_BLOCKING_WAIT'] = '1'
print(f"[{rank}] GPU Configuration:") print(f"[{rank}] env ready")
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( # STEP 3: INIT PROCESS GROUP
nn.Linear(10, 128), # =========================
nn.ReLU(), dist.init_process_group(
nn.Linear(128, 128), backend="nccl",
nn.ReLU(), init_method="env://",
nn.Linear(128, 10) timeout=timedelta(minutes=10),
).to(device) rank=rank,
world_size=world_size
)
ddp_model = nn.parallel.DistributedDataParallel( print(f"[{rank}] process group OK")
model,
device_ids=[0],
output_device=0
)
criterion = nn.MSELoss() # =========================
optimizer = optim.SGD(ddp_model.parameters(), lr=0.001, momentum=0.9) # STEP 4: GPU
# =========================
device = torch.device("cuda:0")
torch.cuda.set_device(device)
print(f"[{rank}] Model initialized with {sum(p.numel() for p in model.parameters())} parameters") model = nn.Sequential(
nn.Linear(10, 128),
nn.ReLU(),
nn.Linear(128, 10)
).to(device)
# STEP 7: Training loop ddp_model = nn.parallel.DistributedDataParallel(
print(f"\n[{rank}] {'='*60}") model,
print(f"[{rank}] Starting Training Loop") device_ids=[0],
print(f"[{rank}] {'='*60}") output_device=0
)
num_epochs = 10 criterion = nn.MSELoss()
batch_size = 32 optimizer = optim.SGD(ddp_model.parameters(), lr=0.001)
num_batches = 5
for epoch in range(num_epochs): # =========================
ddp_model.train() # STEP 5: TRAIN LOOP
epoch_loss = 0.0 # =========================
for epoch in range(5):
ddp_model.train()
loss_sum = 0
for batch_idx in range(num_batches): for _ in range(5):
torch.manual_seed(epoch * num_batches + batch_idx) x = torch.randn(32, 10).to(device)
inputs = torch.randn(batch_size, 10).to(device) y = torch.randn(32, 10).to(device)
labels = torch.randn(batch_size, 10).to(device)
optimizer.zero_grad() optimizer.zero_grad()
outputs = ddp_model(inputs) out = ddp_model(x)
loss = criterion(outputs, labels) loss = criterion(out, y)
loss.backward() loss.backward()
optimizer.step() optimizer.step()
epoch_loss += loss.item() loss_sum += loss.item()
avg_loss = epoch_loss / num_batches print(f"[{rank}] epoch {epoch} loss {loss_sum/5:.4f}")
# Synchronize loss across ranks dist.destroy_process_group()
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}] DONE")
print(f"[{rank}] Epoch {epoch+1}/{num_epochs} | Global Avg Loss: {global_avg_loss:.6f}")
print(f"\n[{rank}] {'='*60}") return {"rank": rank, "status": "success"}
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(): def cleanup_sync_state_func():
"""Clean up synchronization state""" try:
try: if os.path.exists(BARRIER_FILE):
Variable.delete('ddp_sync_state') os.remove(BARRIER_FILE)
print("Sync state cleaned up") except:
except: pass
pass print("Barrier cleaned")
return True return True
def training_summary_func(**context): def training_summary_func(**context):
"""Aggregate and display training results""" ti = context['ti']
ti = context['ti']
print(f"\n{'='*60}") print("\n=== SUMMARY ===")
print(f"DISTRIBUTED TRAINING SUMMARY")
print(f"{'='*60}")
for rank in range(WORLD_SIZE): for rank in range(WORLD_SIZE):
try: res = ti.xcom_pull(task_ids=f"train_rank_{rank}")
result = ti.xcom_pull(task_ids=f"train_rank_{rank}") print(f"Rank {rank}: {res}")
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": "done"}
return {'status': 'success', 'workers': WORLD_SIZE}
# =========================
# DAG
# =========================
with DAG( with DAG(
dag_id='pytorch_distributed_training_ddp_production', dag_id='pytorch_ddp_airflow_fixed',
default_args=default_args, default_args=default_args,
schedule=None, schedule=None,
start_date=pendulum.today('UTC').add(days=-1), start_date=pendulum.today('UTC').add(days=-1),
catchup=False, catchup=False,
tags=['gpu', 'ml', 'distributed'], max_active_runs=1,
max_active_runs=1, tags=['gpu', 'ddp', 'fixed'],
max_active_tasks=10,
) as dag: ) as dag:
# Preparation task prep = PythonOperator(
prep = PythonOperator( task_id='prepare_training',
task_id='prepare_training', python_callable=prepare_training_func,
python_callable=prepare_training_func, queue='gpu',
queue='gpu' )
)
# Training tasks - one per rank training_tasks = []
training_tasks = [] for i in range(WORLD_SIZE):
for i in range(WORLD_SIZE): t = PythonOperator(
task = PythonOperator( task_id=f'train_rank_{i}',
task_id=f'train_rank_{i}', python_callable=run_training_node_func,
python_callable=run_training_node_func, op_kwargs={'rank': i, 'world_size': WORLD_SIZE},
op_kwargs={'rank': i, 'world_size': WORLD_SIZE}, queue='gpu',
queue='gpu', )
) training_tasks.append(t)
training_tasks.append(task)
# Cleanup task cleanup = PythonOperator(
cleanup = PythonOperator( task_id='cleanup',
task_id='cleanup_sync_state', python_callable=cleanup_sync_state_func,
python_callable=cleanup_sync_state_func, queue='gpu',
queue='gpu', trigger_rule='all_done',
trigger_rule='all_done' )
)
# Summary task summary = PythonOperator(
summary = PythonOperator( task_id='summary',
task_id='training_summary', python_callable=training_summary_func,
python_callable=training_summary_func, trigger_rule='all_done',
queue='gpu', )
trigger_rule='all_done'
)
# Set dependencies
prep >> training_tasks >> cleanup >> summary
prep >> training_tasks >> cleanup >> summary