import torch, torchvision import time def get_model(): model = torchvision.models.resnet50() model = model.eval() return model def get_input(batch_size): batch = torch.randn(batch_size, 3, 224, 224) return batch def get_inference_fn(model): def infer_fn(batch): with torch.inference_mode(): output = model(batch) return output return infer_fn def benchmark(infer_fn, batch): # warm-up for _ in range(10): _ = infer_fn(batch) iters = 100 start = time.time() for _ in range(iters): _ = infer_fn(batch) end = time.time() return (end - start) / iters batch_size = 1 model = get_model() batch = get_input(batch_size) infer_fn = get_inference_fn(model) avg_time = benchmark(infer_fn, batch) print(f"\nAverage samples per second: {(batch_size/avg_time):.2f}")