update dags
This commit is contained in:
@@ -1,12 +1,26 @@
|
||||
from airflow import DAG
|
||||
from airflow.operators.python import PythonOperator
|
||||
import pendulum
|
||||
import subprocess
|
||||
from datetime import timedelta
|
||||
|
||||
|
||||
WORLD_SIZE = 2
|
||||
MASTER_ADDR = "127.0.0.1"
|
||||
MASTER_PORT = "29500"
|
||||
|
||||
|
||||
def run_ddp_job():
|
||||
cmd = [
|
||||
"torchrun",
|
||||
"--nproc_per_node=2",
|
||||
"--standalone",
|
||||
"train.py"
|
||||
]
|
||||
|
||||
print("Running:", " ".join(cmd))
|
||||
|
||||
subprocess.run(cmd, check=True)
|
||||
|
||||
return {"status": "success"}
|
||||
|
||||
|
||||
default_args = {
|
||||
@@ -16,77 +30,14 @@ default_args = {
|
||||
}
|
||||
|
||||
|
||||
# =========================
|
||||
# SINGLE SAFE DDP LAUNCHER
|
||||
# =========================
|
||||
def run_ddp_job():
|
||||
import os
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.multiprocessing as mp
|
||||
import torch.nn as nn
|
||||
import torch.optim as optim
|
||||
|
||||
def worker(rank, world_size):
|
||||
os.environ["MASTER_ADDR"] = MASTER_ADDR
|
||||
os.environ["MASTER_PORT"] = MASTER_PORT
|
||||
os.environ["WORLD_SIZE"] = str(world_size)
|
||||
os.environ["RANK"] = str(rank)
|
||||
|
||||
# optional debug
|
||||
os.environ["NCCL_DEBUG"] = "INFO"
|
||||
os.environ["NCCL_ASYNC_ERROR_HANDLING"] = "1"
|
||||
|
||||
torch.cuda.set_device(0)
|
||||
|
||||
dist.init_process_group(
|
||||
backend="nccl",
|
||||
init_method="env://",
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
)
|
||||
|
||||
model = nn.Linear(10, 10).cuda()
|
||||
|
||||
ddp_model = torch.nn.parallel.DistributedDataParallel(
|
||||
model,
|
||||
device_ids=[0]
|
||||
)
|
||||
|
||||
loss_fn = nn.MSELoss()
|
||||
opt = optim.SGD(ddp_model.parameters(), lr=0.01)
|
||||
|
||||
for epoch in range(5):
|
||||
x = torch.randn(32, 10).cuda()
|
||||
y = torch.randn(32, 10).cuda()
|
||||
|
||||
opt.zero_grad()
|
||||
out = ddp_model(x)
|
||||
loss = loss_fn(out, y)
|
||||
loss.backward()
|
||||
opt.step()
|
||||
|
||||
print(f"rank {rank} epoch {epoch} loss {loss.item()}")
|
||||
|
||||
dist.destroy_process_group()
|
||||
|
||||
# IMPORTANT: THIS FIXES NCCL HANG
|
||||
mp.spawn(worker, args=(WORLD_SIZE,), nprocs=WORLD_SIZE, join=True)
|
||||
|
||||
return {"status": "success"}
|
||||
|
||||
|
||||
# =========================
|
||||
# AIRFLOW DAG
|
||||
# =========================
|
||||
with DAG(
|
||||
dag_id="pytorch_ddp_airflow_fixed_stable",
|
||||
dag_id="pytorch_ddp_airflow_fixed_production",
|
||||
default_args=default_args,
|
||||
schedule=None,
|
||||
start_date=pendulum.today("UTC").add(days=-1),
|
||||
catchup=False,
|
||||
max_active_runs=1,
|
||||
tags=["ddp", "gpu", "stable"],
|
||||
tags=["ddp", "torchrun", "stable"],
|
||||
) as dag:
|
||||
|
||||
train = PythonOperator(
|
||||
|
||||
Reference in New Issue
Block a user