Redirect NCCL stderr to stdout, add verbose per-batch logging

- Redirect stderr to stdout via dup2 so Airflow captures NCCL TRACE
  output for all ranks (not just rank 0)
- Log per-batch loss for all ranks
- Log per-epoch summary with local/global avg loss and timing
- Log total training time

Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
2026-04-16 16:24:59 +03:00
parent 0a43170680
commit fdb3e590a6

View File

@@ -38,6 +38,10 @@ def run_training_node_func(rank, world_size):
import socket import socket
from datetime import datetime from datetime import datetime
import time import time
import sys
# Redirect NCCL stderr to stdout so Airflow captures it
os.dup2(sys.stdout.fileno(), sys.stderr.fileno())
# CRITICAL: Set NCCL environment variables BEFORE importing torch # CRITICAL: Set NCCL environment variables BEFORE importing torch
os.environ['NCCL_DEBUG'] = 'TRACE' os.environ['NCCL_DEBUG'] = 'TRACE'
@@ -262,9 +266,12 @@ def run_training_node_func(rank, world_size):
batch_size = 32 batch_size = 32
num_batches = 5 num_batches = 5
training_start = time.time()
for epoch in range(num_epochs): for epoch in range(num_epochs):
ddp_model.train() ddp_model.train()
epoch_loss = 0.0 epoch_loss = 0.0
epoch_start = time.time()
for batch_idx in range(num_batches): for batch_idx in range(num_batches):
torch.manual_seed(epoch * num_batches + batch_idx) torch.manual_seed(epoch * num_batches + batch_idx)
@@ -278,22 +285,24 @@ def run_training_node_func(rank, world_size):
optimizer.step() optimizer.step()
epoch_loss += loss.item() epoch_loss += loss.item()
print(f"[{rank}] Epoch {epoch+1}/{num_epochs} Batch {batch_idx+1}/{num_batches} | Loss: {loss.item():.6f}")
avg_loss = epoch_loss / num_batches avg_loss = epoch_loss / num_batches
epoch_time = time.time() - epoch_start
# Synchronize loss across ranks # Synchronize loss across ranks
loss_tensor = torch.tensor([avg_loss]).to(device) loss_tensor = torch.tensor([avg_loss]).to(device)
dist.all_reduce(loss_tensor, op=dist.ReduceOp.AVG) dist.all_reduce(loss_tensor, op=dist.ReduceOp.AVG)
global_avg_loss = loss_tensor.item() global_avg_loss = loss_tensor.item()
if rank == 0: print(f"[{rank}] Epoch {epoch+1}/{num_epochs} DONE | Local Avg: {avg_loss:.6f} | Global Avg: {global_avg_loss:.6f} | Time: {epoch_time:.3f}s")
print(f"[{rank}] Epoch {epoch+1}/{num_epochs} | Global Avg Loss: {global_avg_loss:.6f}")
total_time = time.time() - training_start
print(f"\n[{rank}] {'='*60}") print(f"\n[{rank}] {'='*60}")
print(f"[{rank}] Training Complete!") print(f"[{rank}] Training Complete!")
print(f"[{rank}] Total training time: {total_time:.3f}s")
print(f"[{rank}] {'='*60}") print(f"[{rank}] {'='*60}")
# STEP 8: Cleanup
dist.destroy_process_group() dist.destroy_process_group()
print(f"[{rank}] Process group destroyed. Finished at {datetime.now()}") print(f"[{rank}] Process group destroyed. Finished at {datetime.now()}")