data_utils.py 2.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102
  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. import numpy as np
  8. mpii_metadata = {
  9. 'layout_name': 'mpii',
  10. 'num_joints': 16,
  11. 'keypoints_symmetry': [
  12. [3, 4, 5, 13, 14, 15],
  13. [0, 1, 2, 10, 11, 12],
  14. ]
  15. }
  16. coco_metadata = {
  17. 'layout_name': 'coco',
  18. 'num_joints': 17,
  19. 'keypoints_symmetry': [
  20. [1, 3, 5, 7, 9, 11, 13, 15],
  21. [2, 4, 6, 8, 10, 12, 14, 16],
  22. ]
  23. }
  24. h36m_metadata = {
  25. 'layout_name': 'h36m',
  26. 'num_joints': 17,
  27. 'keypoints_symmetry': [
  28. [4, 5, 6, 11, 12, 13],
  29. [1, 2, 3, 14, 15, 16],
  30. ]
  31. }
  32. humaneva15_metadata = {
  33. 'layout_name': 'humaneva15',
  34. 'num_joints': 15,
  35. 'keypoints_symmetry': [
  36. [2, 3, 4, 8, 9, 10],
  37. [5, 6, 7, 11, 12, 13]
  38. ]
  39. }
  40. humaneva20_metadata = {
  41. 'layout_name': 'humaneva20',
  42. 'num_joints': 20,
  43. 'keypoints_symmetry': [
  44. [3, 4, 5, 6, 11, 12, 13, 14],
  45. [7, 8, 9, 10, 15, 16, 17, 18]
  46. ]
  47. }
  48. def suggest_metadata(name):
  49. names = []
  50. for metadata in [mpii_metadata, coco_metadata, h36m_metadata, humaneva15_metadata, humaneva20_metadata]:
  51. if metadata['layout_name'] in name:
  52. return metadata
  53. names.append(metadata['layout_name'])
  54. raise KeyError('Cannot infer keypoint layout from name "{}". Tried {}.'.format(name, names))
  55. def import_detectron_poses(path):
  56. # Latin1 encoding because Detectron runs on Python 2.7
  57. data = np.load(path, encoding='latin1')
  58. kp = data['keypoints']
  59. bb = data['boxes']
  60. results = []
  61. for i in range(len(bb)):
  62. if len(bb[i][1]) == 0:
  63. assert i > 0
  64. # Use last pose in case of detection failure
  65. results.append(results[-1])
  66. continue
  67. best_match = np.argmax(bb[i][1][:, 4])
  68. keypoints = kp[i][1][best_match].T.copy()
  69. results.append(keypoints)
  70. results = np.array(results)
  71. return results[:, :, 4:6] # Soft-argmax
  72. #return results[:, :, [0, 1, 3]] # Argmax + score
  73. def import_cpn_poses(path):
  74. data = np.load(path)
  75. kp = data['keypoints']
  76. return kp[:, :, :2]
  77. def import_sh_poses(path):
  78. import h5py
  79. with h5py.File(path) as hf:
  80. positions = hf['poses'].value
  81. return positions.astype('float32')
  82. def suggest_pose_importer(name):
  83. if 'detectron' in name:
  84. return import_detectron_poses
  85. if 'cpn' in name:
  86. return import_cpn_poses
  87. if 'sh' in name:
  88. return import_sh_poses
  89. raise KeyError('Cannot infer keypoint format from name "{}". Tried detectron, cpn, sh.'.format(name))