input_batch = torch.randn((BATCH_SIZE, INPUT_SAMPLES, FEATURE_DIM), device='cuda') labels = torch.randint(0, 2, (BATCH_SIZE, INPUT_SAMPLES), device='cuda', dtype=torch.int64) benchmark(batched_sample_data, input_batch, labels) benchmark(opt_sample_data, input_batch, labels) benchmark(torch.compile(opt_sample_data), input_batch, labels)