class DataPrefetcher: def __init__(self, loader): self.loader = iter(loader) self.stream = torch.cuda.Stream() self.next_batch = None self.preload() def preload(self): try: data, targets = next(self.loader) with torch.cuda.stream(self.stream): with nvtx.annotate("copy batch", color="yellow"): next_data = data.to(DEVICE, non_blocking=True) next_targets = targets.to(DEVICE, non_blocking=True) self.next_batch = (next_data, next_targets) except: self.next_batch = (None, None) def __iter__(self): return self def __next__(self): torch.cuda.current_stream().wait_stream(self.stream) data, targets = self.next_batch self.preload() return data, targets data_iter = DataPrefetcher(train_loader) for i in range(TOTAL_STEPS): if i == WARMUP_STEPS: torch.cuda.synchronize() start_time = time.perf_counter() profiler.start() elif i == WARMUP_STEPS + PROFILE_STEPS: torch.cuda.synchronize() profiler.stop() end_time = time.perf_counter() with nvtx.annotate(f"Batch {i}", color="blue"): with nvtx.annotate("get batch", color="red"): batch = next(data_iter) with nvtx.annotate("Compute", color="green"): loss = compute_step(model, batch, optimizer) total_time = end_time - start_time throughput = PROFILE_STEPS / total_time print(f"Throughput: {throughput:.2f} steps/sec")