update dags
This commit is contained in:
@@ -1,12 +1,26 @@
|
|||||||
from airflow import DAG
|
from airflow import DAG
|
||||||
from airflow.operators.python import PythonOperator
|
from airflow.operators.python import PythonOperator
|
||||||
import pendulum
|
import pendulum
|
||||||
|
import subprocess
|
||||||
from datetime import timedelta
|
from datetime import timedelta
|
||||||
|
|
||||||
|
|
||||||
WORLD_SIZE = 2
|
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 = {
|
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(
|
with DAG(
|
||||||
dag_id="pytorch_ddp_airflow_fixed_stable",
|
dag_id="pytorch_ddp_airflow_fixed_production",
|
||||||
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=["ddp", "gpu", "stable"],
|
tags=["ddp", "torchrun", "stable"],
|
||||||
) as dag:
|
) as dag:
|
||||||
|
|
||||||
train = PythonOperator(
|
train = PythonOperator(
|
||||||
|
|||||||
35
dags/train.py
Normal file
35
dags/train.py
Normal file
@@ -0,0 +1,35 @@
|
|||||||
|
import torch
|
||||||
|
import torch.distributed as dist
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.optim as optim
|
||||||
|
|
||||||
|
|
||||||
|
def main():
|
||||||
|
dist.init_process_group("nccl")
|
||||||
|
|
||||||
|
rank = dist.get_rank()
|
||||||
|
torch.cuda.set_device(0)
|
||||||
|
|
||||||
|
model = nn.Linear(10, 10).cuda()
|
||||||
|
ddp = torch.nn.parallel.DistributedDataParallel(model, device_ids=[0])
|
||||||
|
|
||||||
|
opt = optim.SGD(ddp.parameters(), lr=0.01)
|
||||||
|
loss_fn = nn.MSELoss()
|
||||||
|
|
||||||
|
for i in range(5):
|
||||||
|
x = torch.randn(32, 10).cuda()
|
||||||
|
y = torch.randn(32, 10).cuda()
|
||||||
|
|
||||||
|
opt.zero_grad()
|
||||||
|
out = ddp(x)
|
||||||
|
loss = loss_fn(out, y)
|
||||||
|
loss.backward()
|
||||||
|
opt.step()
|
||||||
|
|
||||||
|
print(f"rank {rank} step {i} loss {loss.item()}")
|
||||||
|
|
||||||
|
dist.destroy_process_group()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
Reference in New Issue
Block a user