From 0b46d00bb2788e8222e249c48147d72ba94811a1 Mon Sep 17 00:00:00 2001 From: George Stykalin Date: Wed, 15 Apr 2026 12:04:00 +0300 Subject: [PATCH] update dags --- dags/test-train-pytorch.py | 85 ++++++++------------------------------ dags/train.py | 35 ++++++++++++++++ 2 files changed, 53 insertions(+), 67 deletions(-) create mode 100644 dags/train.py diff --git a/dags/test-train-pytorch.py b/dags/test-train-pytorch.py index 8f2e16f..607df4b 100644 --- a/dags/test-train-pytorch.py +++ b/dags/test-train-pytorch.py @@ -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( diff --git a/dags/train.py b/dags/train.py new file mode 100644 index 0000000..5f898a0 --- /dev/null +++ b/dags/train.py @@ -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()