def export_to_onnx(model, onnx_path="resnet50.onnx"): dummy_input = torch.randn(4, 3, 224, 224) batch = torch.export.Dim("batch") torch.onnx.export( model, dummy_input, onnx_path, input_names=["input"], output_names=["output"], dynamic_shapes=((batch, torch.export.Dim.STATIC, torch.export.Dim.STATIC, torch.export.Dim.STATIC), ), dynamo=True ) return onnx_path def onnx_infer_fn(onnx_path): import onnxruntime as ort sess = ort.InferenceSession( onnx_path, providers=["CPUExecutionProvider"] ) input_name = sess.get_inputs()[0].name def infer_fn(batch): result = sess.run(None, {input_name: batch}) return result return infer_fn batch_size = 8 model = get_model() onnx_path = export_to_onnx(model) batch = get_input(batch_size).numpy() infer_fn = onnx_infer_fn(onnx_path) avg_time = benchmark(infer_fn, batch) print(f"\nAverage samples per second: {(batch_size/avg_time):.2f}")