speed.py 1.7 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849
  1. from models.experimental import attempt_load
  2. from torch2trt.trt_model import TrtModel
  3. import argparse
  4. import torch
  5. import time
  6. from tqdm import tqdm
  7. def run(model,img,warmup_iter,iter):
  8. print('start warm up...')
  9. for _ in tqdm(range(warmup_iter)):
  10. model(img)
  11. print('start calculate...')
  12. torch.cuda.synchronize()
  13. start = time.time()
  14. for __ in tqdm(range(iter)):
  15. model(img)
  16. torch.cuda.synchronize()
  17. end = time.time()
  18. return ((end - start) * 1000)/float(iter)
  19. if __name__ == '__main__':
  20. parser = argparse.ArgumentParser()
  21. parser.add_argument('--torch_path', type=str,required=True, help='torch weights path')
  22. parser.add_argument('--trt_path', type=str,required=True, help='tensorrt weights path')
  23. parser.add_argument('--device', type=int,default=0, help='cuda device')
  24. parser.add_argument('--img_shape', type=list,default=[1,3,640,640], help='tensorrt weights path')
  25. parser.add_argument('--warmup_iter', type=int, default=100,help='warm up iter')
  26. parser.add_argument('--iter', type=int, default=300,help='average elapsed time of iterations')
  27. opt = parser.parse_args()
  28. # -----------------------torch-----------------------------------------
  29. img = torch.zeros(opt.img_shape)
  30. model = attempt_load(opt.torch_path, map_location=torch.device('cpu')) # load FP32 model
  31. model.eval()
  32. total_time=run(model.to(opt.device),img.to(opt.device),opt.warmup_iter,opt.iter)
  33. print('Pytorch is %.2f ms/img'%total_time)
  34. # -----------------------tensorrt-----------------------------------------
  35. model=TrtModel(opt.trt_path)
  36. total_time=run(model,img.numpy(),opt.warmup_iter,opt.iter)
  37. model.destroy()
  38. print('TensorRT is %.2f ms/img'%total_time)