def train(use_ddp=False): # detect the env vars set by torchrun rank = int(os.environ.get("RANK", 0)) local_rank = int(os.environ.get("LOCAL_RANK", 0)) torch.cuda.set_device(local_rank) model = get_model().to(local_rank) criterion = nn.CrossEntropyLoss().to(local_rank) if use_ddp: model = configure_ddp(model, rank) optimizer = optim.SGD(model.parameters()) data_iter = get_data_iter(rank, WORLD_SIZE) model.train() for i in range(TOTAL_STEPS): # Schedule Profiling if i == WARMUP_STEPS: torch.cuda.synchronize() start_time = time.perf_counter() torch.cuda.profiler.start() elif i == WARMUP_STEPS + PROFILE_STEPS: torch.cuda.synchronize() torch.cuda.profiler.stop() end_time = time.perf_counter() with nvtx.annotate(f"Batch {i}", color="blue"): with nvtx.annotate("get batch", color="red"): data, target = next(data_iter) data = data.to(local_rank, non_blocking=True) target = target.to(local_rank, non_blocking=True) with nvtx.annotate("forward", color="green"): output = model(data) loss = criterion(output, target) with nvtx.annotate("backward", color="purple"): optimizer.zero_grad() loss.backward() with nvtx.annotate("optimizer step", color="yellow"): optimizer.step() if use_ddp: dist.destroy_process_group() if rank == 0: total_time = end_time - start_time print(f"Throughput: {PROFILE_STEPS/total_time:.2f} steps/sec") if __name__ == "__main__": # enable ddp if run with torchrun train(use_ddp="RANK" in os.environ)