update dags

This commit is contained in:
2026-04-15 12:04:00 +03:00
parent 9b365092b5
commit 0b46d00bb2
2 changed files with 53 additions and 67 deletions

View File

@@ -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(

35
dags/train.py Normal file
View 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()