Files
test-dags/dags/test-train-pytorch.py
2026-04-15 00:10:17 +03:00

135 lines
5.3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

from airflow import DAG
from airflow.operators.python import PythonOperator
import pendulum
from datetime import timedelta
WORLD_SIZE = 2
# Адрес мастера — headless DNS пода воркера
MASTER_ADDR = "airflow-worker-gpu-0.airflow-worker-gpu"
MASTER_PORT = "29500"
default_args = {
"owner": "airflow",
"retries": 0,
"execution_timeout": timedelta(hours=1),
}
def run_training_node_func(rank: int, world_size: int):
import os
import time
import socket
import torch
import torch.distributed as dist
import torch.nn as nn
import torch.optim as optim
from datetime import datetime
print("=" * 80)
print(f"RANK {rank}/{world_size} HOST={socket.gethostname()} {datetime.now()}")
print("=" * 80)
# ── статический барьер: ждём пока оба воркера подтянутся ──────────────────
time.sleep(10)
# ── переменные окружения для NCCL + InfiniBand ────────────────────────────
os.environ.update(
{
"MASTER_ADDR": MASTER_ADDR,
"MASTER_PORT": MASTER_PORT,
"WORLD_SIZE": str(world_size),
"RANK": str(rank),
# InfiniBand: указываем IPoIB-интерфейс вместо eth0
# Разрешаем NCCL использовать RDMA (IB verbs)
"NCCL_IB_DISABLE": "0",
# GPUDirect RDMA — если драйвер поддерживает
"NCCL_P2P_DISABLE": "0",
# Явно форсируем IB transport (опционально, NCCL сам выберет,
# но полезно для отладки)
"NCCL_NET": "IB",
# Отладка: INFO покажет какой транспорт выбран,
# поменяй на TRACE для полного вывода
"NCCL_DEBUG": "INFO",
"NCCL_DEBUG_SUBSYS": "NET,INIT",
"TORCH_NCCL_BLOCKING_WAIT": "1",
}
)
print(f"[{rank}] MASTER={MASTER_ADDR}:{MASTER_PORT} IB_IF=ib0")
# ── init process group ────────────────────────────────────────────────────
dist.init_process_group(
backend="nccl",
init_method="env://",
rank=rank,
world_size=world_size,
timeout=timedelta(minutes=5),
)
print(f"[{rank}] dist.init_process_group OK")
# ── GPU binding ───────────────────────────────────────────────────────────
torch.cuda.set_device(0)
device = torch.device("cuda:0")
print(f"[{rank}] GPU={torch.cuda.get_device_name(0)}")
# ── минимальная модель ────────────────────────────────────────────────────
model = nn.Sequential(
nn.Linear(10, 128),
nn.ReLU(),
nn.Linear(128, 10),
).to(device)
ddp_model = nn.parallel.DistributedDataParallel(model, device_ids=[0])
optimizer = optim.SGD(ddp_model.parameters(), lr=0.001)
loss_fn = nn.MSELoss()
# ── allreduce smoke-test перед обучением ──────────────────────────────────
probe = torch.ones(1).to(device) * rank
dist.all_reduce(probe, op=dist.ReduceOp.SUM)
expected = sum(range(world_size))
assert probe.item() == expected, f"allreduce mismatch: got {probe.item()}"
print(f"[{rank}] allreduce smoke-test PASSED (sum={probe.item()})")
# ── train loop ────────────────────────────────────────────────────────────
for epoch in range(5):
x = torch.randn(32, 10).to(device)
y = torch.randn(32, 10).to(device)
optimizer.zero_grad()
out = ddp_model(x)
loss = loss_fn(out, y)
loss.backward()
optimizer.step()
print(f"[{rank}] epoch={epoch} loss={loss.item():.4f}")
dist.destroy_process_group()
print(f"[{rank}] DONE")
return {"rank": rank, "status": "ok"}
# ── DAG ───────────────────────────────────────────────────────────────────────
with DAG(
dag_id="ddp_ib_test",
start_date=pendulum.today("UTC").add(days=-1),
schedule=None,
catchup=False,
max_active_runs=1,
max_active_tasks=WORLD_SIZE,
default_args=default_args,
) as dag:
tasks = []
for r in range(WORLD_SIZE):
t = PythonOperator(
task_id=f"train_rank_{r}",
python_callable=run_training_node_func,
op_kwargs={"rank": r, "world_size": WORLD_SIZE},
queue="gpu",
)
tasks.append(t)
# запускаем параллельно — без зависимостей между рангами
# (prep/cleanup убраны, это минимальный тест)