def copy_to_device(batch, device): return { key: val.to(device=device, non_blocking=True) for key, val in batch.items() } def train_step(model, device, optimizer, batch): # copy data to device batch = copy_to_device(batch, device) optimizer.zero_grad() with torch.amp.autocast('cuda', dtype=torch.bfloat16): outputs = model(**batch) loss = outputs.loss loss.backward() optimizer.step() return loss def train(local_rank=0, world_size=1, compile=False): # specify log settings torch._logging.set_logs( graph_breaks=True, recompiles=True, perf_hints=True ) torch.cuda.set_device(local_rank) device = torch.cuda.current_device() if world_size > 1: # DDP setup import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP os.environ['MASTER_ADDR'] = 'localhost' os.environ['MASTER_PORT'] = str(2222) dist.init_process_group('nccl', rank=local_rank, world_size=world_size) # configure pad_to_longest and optional alignment dataloader = get_dataloader(pad_to_longest=False, align=None) model = get_model() model = model.to(device) if world_size > 1: model = DDP(model, [local_rank]) optimizer = torch.optim.Adam(model.parameters()) if compile: # uncomment to run pre-compile warmup - required for some optimizations # batch = next(iter(dataloader)) # train_step(model, device, optimizer, batch) model, optimizer = apply_compilation(model, optimizer) warmup = 20 active = 100 total_steps = warmup + active t0 = time.perf_counter() for idx, batch in enumerate(dataloader, start=1): # apply train step train_step(model, device, optimizer, batch) if idx == warmup: torch.cuda.synchronize() print(f'warmup time: {time.perf_counter()-t0}') t0 = time.perf_counter() elif idx == total_steps: break if local_rank == 0: torch.cuda.synchronize() total_time = time.perf_counter() - t0 print(f'average throughput: {active / total_time}') if world_size > 1: dist.destroy_process_group() if __name__ == '__main__': # specify inductor cache dir inductor_cache_dir = '/tmp/inductor_cache' os.environ['TORCHINDUCTOR_CACHE_DIR'] = inductor_cache_dir # clean up compiler cache torch._dynamo.reset() shutil.rmtree(inductor_cache_dir, ignore_errors=True) world_size = 1 torch.multiprocessing.spawn( fn=train, args=(world_size,), nprocs=world_size, join=True )