update dags

This commit is contained in:
2026-04-15 12:02:41 +03:00
parent 5f7db605a6
commit 9b365092b5

View File

@@ -2,219 +2,95 @@ from airflow import DAG
from airflow.operators.python import PythonOperator from airflow.operators.python import PythonOperator
import pendulum import pendulum
from datetime import timedelta from datetime import timedelta
import os
import time
# --- CONFIGURATION ---
WORLD_SIZE = 2 WORLD_SIZE = 2
MASTER_ADDR = "airflow-worker-gpu-0.airflow-worker-gpu" MASTER_ADDR = "127.0.0.1"
MASTER_PORT = "29500" MASTER_PORT = "29500"
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),
} }
def prepare_training_func(): # =========================
# reset barrier # SINGLE SAFE DDP LAUNCHER
try: # =========================
if os.path.exists(BARRIER_FILE): def run_ddp_job():
os.remove(BARRIER_FILE) import os
except:
pass
print("Barrier initialized")
return True
def run_training_node_func(rank, world_size):
import socket
import torch import torch
import torch.distributed as dist import torch.distributed as dist
import torch.multiprocessing as mp
import torch.nn as nn import torch.nn as nn
import torch.optim as optim import torch.optim as optim
from datetime import datetime
print("=" * 60) def worker(rank, world_size):
print(f"Rank {rank}/{world_size} START {datetime.now()}") os.environ["MASTER_ADDR"] = MASTER_ADDR
print(f"Hostname: {socket.gethostname()}") os.environ["MASTER_PORT"] = MASTER_PORT
print("=" * 60) os.environ["WORLD_SIZE"] = str(world_size)
os.environ["RANK"] = str(rank)
# ========================= # optional debug
# STEP 1: SAFE BARRIER os.environ["NCCL_DEBUG"] = "INFO"
# ========================= os.environ["NCCL_ASYNC_ERROR_HANDLING"] = "1"
print(f"[{rank}] entering barrier sync...")
with open(BARRIER_FILE, "a+") as f: torch.cuda.set_device(0)
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( dist.init_process_group(
backend="nccl", backend="nccl",
init_method="env://", init_method="env://",
timeout=timedelta(minutes=10),
rank=rank, rank=rank,
world_size=world_size world_size=world_size,
) )
print(f"[{rank}] process group OK") model = nn.Linear(10, 10).cuda()
# ========================= ddp_model = torch.nn.parallel.DistributedDataParallel(
# 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, model,
device_ids=[0], device_ids=[0]
output_device=0
) )
criterion = nn.MSELoss() loss_fn = nn.MSELoss()
optimizer = optim.SGD(ddp_model.parameters(), lr=0.001) opt = optim.SGD(ddp_model.parameters(), lr=0.01)
# =========================
# STEP 5: TRAIN LOOP
# =========================
for epoch in range(5): for epoch in range(5):
ddp_model.train() x = torch.randn(32, 10).cuda()
loss_sum = 0 y = torch.randn(32, 10).cuda()
for _ in range(5): opt.zero_grad()
x = torch.randn(32, 10).to(device)
y = torch.randn(32, 10).to(device)
optimizer.zero_grad()
out = ddp_model(x) out = ddp_model(x)
loss = criterion(out, y) loss = loss_fn(out, y)
loss.backward() loss.backward()
optimizer.step() opt.step()
loss_sum += loss.item() print(f"rank {rank} epoch {epoch} loss {loss.item()}")
print(f"[{rank}] epoch {epoch} loss {loss_sum/5:.4f}")
dist.destroy_process_group() dist.destroy_process_group()
print(f"[{rank}] DONE") # IMPORTANT: THIS FIXES NCCL HANG
mp.spawn(worker, args=(WORLD_SIZE,), nprocs=WORLD_SIZE, join=True)
return {"rank": rank, "status": "success"} return {"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 # AIRFLOW DAG
# ========================= # =========================
with DAG( with DAG(
dag_id='pytorch_ddp_airflow_fixed', dag_id="pytorch_ddp_airflow_fixed_stable",
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,
max_active_runs=1, max_active_runs=1,
tags=['gpu', 'ddp', 'fixed'], tags=["ddp", "gpu", "stable"],
) as dag: ) as dag:
prep = PythonOperator( train = PythonOperator(
task_id='prepare_training', task_id="train_ddp",
python_callable=prepare_training_func, python_callable=run_ddp_job,
queue='gpu', 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