Files
test-dags/dags/test-train-pytorch.py
2026-04-15 11:54:13 +03:00

221 lines
5.2 KiB
Python

from airflow import DAG
from airflow.operators.python import PythonOperator
import pendulum
from datetime import timedelta
import os
import time
# --- CONFIGURATION ---
WORLD_SIZE = 2
MASTER_ADDR = "airflow-worker-gpu-0.airflow-worker-gpu"
MASTER_PORT = "29500"
NCCL_TIMEOUT = 1800
BARRIER_FILE = f"/tmp/ddp_barrier_{WORLD_SIZE}"
default_args = {
'owner': 'airflow',
'retries': 1,
'retry_delay': timedelta(minutes=2),
'execution_timeout': timedelta(hours=2),
}
def prepare_training_func():
# reset barrier
try:
if os.path.exists(BARRIER_FILE):
os.remove(BARRIER_FILE)
except:
pass
print("Barrier initialized")
return True
def run_training_node_func(rank, world_size):
import socket
import torch
import torch.distributed as dist
import torch.nn as nn
import torch.optim as optim
from datetime import datetime
print("=" * 60)
print(f"Rank {rank}/{world_size} START {datetime.now()}")
print(f"Hostname: {socket.gethostname()}")
print("=" * 60)
# =========================
# STEP 1: SAFE BARRIER
# =========================
print(f"[{rank}] entering barrier sync...")
with open(BARRIER_FILE, "a+") as f:
f.write(f"{rank}\n")
# wait for all ranks
while True:
try:
with open(BARRIER_FILE, "r") as f:
ready = set(f.read().strip().splitlines())
print(f"[{rank}] ready workers: {ready}")
if len(ready) == world_size:
print(f"[{rank}] ALL WORKERS READY")
break
except FileNotFoundError:
pass
time.sleep(1)
time.sleep(2) # small stabilization delay
# =========================
# STEP 2: NCCL ENV
# =========================
os.environ['MASTER_ADDR'] = MASTER_ADDR
os.environ['MASTER_PORT'] = MASTER_PORT
os.environ['WORLD_SIZE'] = str(world_size)
os.environ['RANK'] = str(rank)
# IMPORTANT FIX (IB stability toggle)
os.environ["NCCL_IB_DISABLE"] = "1" # <- FIX #2 (can turn OFF later)
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}] env ready")
# =========================
# STEP 3: INIT PROCESS GROUP
# =========================
dist.init_process_group(
backend="nccl",
init_method="env://",
timeout=timedelta(minutes=10),
rank=rank,
world_size=world_size
)
print(f"[{rank}] process group OK")
# =========================
# STEP 4: GPU
# =========================
device = torch.device("cuda:0")
torch.cuda.set_device(device)
model = nn.Sequential(
nn.Linear(10, 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)
# =========================
# STEP 5: TRAIN LOOP
# =========================
for epoch in range(5):
ddp_model.train()
loss_sum = 0
for _ in range(5):
x = torch.randn(32, 10).to(device)
y = torch.randn(32, 10).to(device)
optimizer.zero_grad()
out = ddp_model(x)
loss = criterion(out, y)
loss.backward()
optimizer.step()
loss_sum += loss.item()
print(f"[{rank}] epoch {epoch} loss {loss_sum/5:.4f}")
dist.destroy_process_group()
print(f"[{rank}] DONE")
return {"rank": rank, "status": "success"}
def cleanup_sync_state_func():
try:
if os.path.exists(BARRIER_FILE):
os.remove(BARRIER_FILE)
except:
pass
print("Barrier cleaned")
return True
def training_summary_func(**context):
ti = context['ti']
print("\n=== SUMMARY ===")
for rank in range(WORLD_SIZE):
res = ti.xcom_pull(task_ids=f"train_rank_{rank}")
print(f"Rank {rank}: {res}")
return {"status": "done"}
# =========================
# DAG
# =========================
with DAG(
dag_id='pytorch_ddp_airflow_fixed',
default_args=default_args,
schedule=None,
start_date=pendulum.today('UTC').add(days=-1),
catchup=False,
max_active_runs=1,
tags=['gpu', 'ddp', 'fixed'],
) as dag:
prep = PythonOperator(
task_id='prepare_training',
python_callable=prepare_training_func,
queue='gpu',
)
training_tasks = []
for i in range(WORLD_SIZE):
t = 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(t)
cleanup = PythonOperator(
task_id='cleanup',
python_callable=cleanup_sync_state_func,
queue='gpu',
trigger_rule='all_done',
)
summary = PythonOperator(
task_id='summary',
python_callable=training_summary_func,
trigger_rule='all_done',
)
prep >> training_tasks >> cleanup >> summary