From 6dddc69d117581ee06145b66d96b1b0d57849f48 Mon Sep 17 00:00:00 2001 From: George Stykalin Date: Wed, 15 Apr 2026 12:10:19 +0300 Subject: [PATCH] update dags --- dags/test-train-pytorch.py | 73 ++++++++++++++++++++------------------ dags/train.py | 31 ++++++++-------- 2 files changed, 54 insertions(+), 50 deletions(-) diff --git a/dags/test-train-pytorch.py b/dags/test-train-pytorch.py index 0e3eeb6..c26197e 100644 --- a/dags/test-train-pytorch.py +++ b/dags/test-train-pytorch.py @@ -1,77 +1,80 @@ from airflow import DAG -from airflow.operators.python import PythonOperator +from airflow.providers.standard.operators.python import PythonOperator from airflow.operators.bash import BashOperator from datetime import datetime import os -# ----------------------------- -# CONFIG -# ----------------------------- -WORLD_SIZE = 2 -MASTER_PORT = 29500 - - default_args = { "owner": "airflow", } -# ----------------------------- -# PREP TASK -# ----------------------------- +# ----------------------- +# PREP +# ----------------------- def prepare_training(): - print("Preparing dataset / env for DDP") + print("Preparing DDP training environment") + os.environ["TOKENIZERS_PARALLELISM"] = "false" - return True + os.environ["NCCL_DEBUG"] = "INFO" + + return "prepared" -# ----------------------------- -# CLEANUP TASK -# ----------------------------- +# ----------------------- +# CLEANUP +# ----------------------- def cleanup(): - print("Cleaning up training artifacts") - return True + print("Cleaning up after training") + return "cleaned" -# ----------------------------- +# ----------------------- # DAG -# ----------------------------- +# ----------------------- with DAG( dag_id="pytorch_ddp_airflow_fixed_production", - default_args=default_args, start_date=datetime(2024, 1, 1), schedule=None, catchup=False, - tags=["ddp", "pytorch", "fixed"], + default_args=default_args, + tags=["ddp", "pytorch", "gpu"], ) as dag: - prepare = PythonOperator( + prepare_training_task = PythonOperator( task_id="prepare_training", python_callable=prepare_training, + queue="gpu", # ✅ FIX #1 ) - # --------------------------------------------------------- - # MAIN FIX: use torchrun instead of spawn inside Python - # --------------------------------------------------------- - train_ddp = BashOperator( + train_ddp_task = BashOperator( task_id="train_ddp", - bash_command=f""" + bash_command=""" + set -e + export NCCL_DEBUG=INFO export NCCL_ASYNC_ERROR_HANDLING=1 + export NCCL_IB_DISABLE=1 - export NCCL_IB_DISABLE=1 # safe default (enable later if needed) + echo "Starting torchrun DDP training" torchrun \ - --nproc_per_node={WORLD_SIZE} \ - --master_port={MASTER_PORT} \ - train.py - """ + --nproc_per_node=2 \ + --master_port=29500 \ + /opt/airflow/dags/repo/train.py + + echo "Training finished" + """, + queue="gpu", # ✅ FIX #1 ) - finish = PythonOperator( + cleanup_task = PythonOperator( task_id="cleanup", python_callable=cleanup, + queue="gpu", # ✅ FIX #1 ) - prepare >> train_ddp >> finish + + # DAG FLOW + prepare_training_task >> train_ddp_task >> cleanup_task diff --git a/dags/train.py b/dags/train.py index 5f898a0..7622e4a 100644 --- a/dags/train.py +++ b/dags/train.py @@ -1,33 +1,34 @@ +import os import torch import torch.distributed as dist -import torch.nn as nn -import torch.optim as optim +import time def main(): dist.init_process_group("nccl") + local_rank = int(os.environ["LOCAL_RANK"]) + torch.cuda.set_device(local_rank) + rank = dist.get_rank() - torch.cuda.set_device(0) + world_size = dist.get_world_size() - model = nn.Linear(10, 10).cuda() - ddp = torch.nn.parallel.DistributedDataParallel(model, device_ids=[0]) + print(f"[rank {rank}/{world_size}] started") - opt = optim.SGD(ddp.parameters(), lr=0.01) - loss_fn = nn.MSELoss() + model = torch.nn.Linear(16, 16).cuda() - for i in range(5): - x = torch.randn(32, 10).cuda() - y = torch.randn(32, 10).cuda() + for step in range(50): + x = torch.randn(32, 16).cuda() + loss = model(x).sum() - opt.zero_grad() - out = ddp(x) - loss = loss_fn(out, y) loss.backward() - opt.step() - print(f"rank {rank} step {i} loss {loss.item()}") + if rank == 0 and step % 10 == 0: + print(f"step={step}, loss={loss.item()}") + time.sleep(0.2) + + dist.barrier() dist.destroy_process_group()