def benchmark(fn, input, labels): # warm-up for _ in range(20): _ = fn(input, labels) iters = 100 start = torch.cuda.Event(enable_timing=True) end = torch.cuda.Event(enable_timing=True) torch.cuda.synchronize() start.record() for _ in range(iters): _ = fn(input, labels) end.record() torch.cuda.synchronize() avg_time = start.elapsed_time(end) / iters print(f"{fn.__name__} average step time: {(avg_time):.4f} ms") benchmark(sample_data, input_samples, labels) benchmark(opt_sample_data, input_samples, labels)