# Set tf mixed precision policy to bfloat16 tf.keras.mixed_precision.set_global_policy('mixed_bfloat16') # Set torch matmul precision to high torch.set_float32_matmul_precision('high') @tf.function def tf_infer_fn(batch): return tf_model(batch, training=False) def get_torch_infer_fn(model): def infer_fn(batch): with torch.inference_mode(), torch.amp.autocast( DEVICE, dtype=torch.bfloat16, enabled=DEVICE=='cuda' ): output = model(batch) return output return infer_fn def benchmark(infer_fn, batch): # warm-up for _ in range(20): _ = infer_fn(batch) start = torch.cuda.Event(enable_timing=True) end = torch.cuda.Event(enable_timing=True) torch.cuda.synchronize() start.record() iters = 100 for _ in range(iters): _ = infer_fn(batch) end.record() torch.cuda.synchronize() return start.elapsed_time(end) / iters # assess throughput of TF model avg_time = benchmark(tf_infer_fn, tf_input) print(f"\nTensorFlow average step time: {(avg_time):.4f}") # assess throughput of converted model torch_infer_fn = get_torch_infer_fn(converted_model) avg_time = benchmark(torch_infer_fn, torch_input) print(f"\nConverted model average step time: {(avg_time):.4f}") # assess throughput of compiled model torch_infer_fn = get_torch_infer_fn(torch.compile(converted_model)) avg_time = benchmark(torch_infer_fn, torch_input) print(f"\nCompiled model average step time: {(avg_time):.4f}") # assess throughput of torch ViT from transformers import ViTForImageClassification torch_model = ViTForImageClassification(vit_config).to(DEVICE) torch_infer_fn = get_torch_infer_fn(torch_model) avg_time = benchmark(torch_infer_fn, torch_input) print(f"\nPyTorch ViT model average step time: {(avg_time):.4f}") # assess throughput of compiled torch ViT torch_infer_fn = get_torch_infer_fn(torch.compile(torch_model)) avg_time = benchmark(torch_infer_fn, torch_input) print(f"\nCompiled ViT model average step time: {(avg_time):.4f}")