infer_video.py 3.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100
  1. # Copyright (c) 2018-present, Facebook, Inc.
  2. # All rights reserved.
  3. #
  4. # This source code is licensed under the license found in the
  5. # LICENSE file in the root directory of this source tree.
  6. #
  7. """Perform inference on a single video or all videos with a certain extension
  8. (e.g., .mp4) in a folder.
  9. """
  10. from infer_simple import *
  11. import subprocess as sp
  12. import numpy as np
  13. def get_resolution(filename):
  14. command = ['ffprobe', '-v', 'error', '-select_streams', 'v:0',
  15. '-show_entries', 'stream=width,height', '-of', 'csv=p=0', filename]
  16. pipe = sp.Popen(command, stdout=sp.PIPE, bufsize=-1)
  17. for line in pipe.stdout:
  18. w, h = line.decode().strip().split(',')
  19. return int(w), int(h)
  20. def read_video(filename):
  21. w, h = get_resolution(filename)
  22. command = ['ffmpeg',
  23. '-i', filename,
  24. '-f', 'image2pipe',
  25. '-pix_fmt', 'bgr24',
  26. '-vsync', '0',
  27. '-vcodec', 'rawvideo', '-']
  28. pipe = sp.Popen(command, stdout=sp.PIPE, bufsize=-1)
  29. while True:
  30. data = pipe.stdout.read(w*h*3)
  31. if not data:
  32. break
  33. yield np.frombuffer(data, dtype='uint8').reshape((h, w, 3))
  34. def main(args):
  35. logger = logging.getLogger(__name__)
  36. merge_cfg_from_file(args.cfg)
  37. cfg.NUM_GPUS = 1
  38. args.weights = cache_url(args.weights, cfg.DOWNLOAD_CACHE)
  39. assert_and_infer_cfg(cache_urls=False)
  40. model = infer_engine.initialize_model_from_cfg(args.weights)
  41. dummy_coco_dataset = dummy_datasets.get_coco_dataset()
  42. if os.path.isdir(args.im_or_folder):
  43. im_list = glob.iglob(args.im_or_folder + '/*.' + args.image_ext)
  44. else:
  45. im_list = [args.im_or_folder]
  46. for video_name in im_list:
  47. out_name = os.path.join(
  48. args.output_dir, os.path.basename(video_name)
  49. )
  50. print('Processing {}'.format(video_name))
  51. boxes = []
  52. segments = []
  53. keypoints = []
  54. for frame_i, im in enumerate(read_video(video_name)):
  55. logger.info('Frame {}'.format(frame_i))
  56. timers = defaultdict(Timer)
  57. t = time.time()
  58. with c2_utils.NamedCudaScope(0):
  59. cls_boxes, cls_segms, cls_keyps = infer_engine.im_detect_all(
  60. model, im, None, timers=timers
  61. )
  62. logger.info('Inference time: {:.3f}s'.format(time.time() - t))
  63. for k, v in timers.items():
  64. logger.info(' | {}: {:.3f}s'.format(k, v.average_time))
  65. boxes.append(cls_boxes)
  66. segments.append(cls_segms)
  67. keypoints.append(cls_keyps)
  68. # Video resolution
  69. metadata = {
  70. 'w': im.shape[1],
  71. 'h': im.shape[0],
  72. }
  73. np.savez_compressed(out_name, boxes=boxes, segments=segments, keypoints=keypoints, metadata=metadata)
  74. if __name__ == '__main__':
  75. workspace.GlobalInit(['caffe2', '--caffe2_log_level=0'])
  76. setup_logging(__name__)
  77. args = parse_args()
  78. main(args)