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