import os, time, random, torch torch.manual_seed(42) random.seed(42) BATCH_SIZE = 64 EXPORT_PATH = '/tmp/export/' def test_inference(model_path=EXPORT_PATH, mode=None, compile=False): device = 'cuda' if torch.cuda.is_available() else 'cpu' rnd_image = torch.randn(BATCH_SIZE, 3, 224, 224).to(device) encoder, decoder = load_model(model_path, mode) encoder = encoder.to(device) decoder = decoder.to(device) if compile: encoder = torch.compile(encoder, mode="reduce-overhead") decoder = torch.compile(decoder, dynamic=True) # run a few warmup rounds for i in range(10): image_to_text_generator(encoder, decoder, random_image) t0 = time.perf_counter() # optionally enable mixed precision with torch.amp.autocast(device, dtype=torch.bfloat16, enabled=True): with torch.no_grad(): caption = image_to_text_generator(encoder, decoder, rnd_image) total_time = time.perf_counter() - t0 print(f'batched inference total time: {total_time}')