run.py 41 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862
  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. from common.arguments import parse_args
  9. import torch
  10. import torch.nn as nn
  11. import torch.nn.functional as F
  12. import torch.optim as optim
  13. import os
  14. import sys
  15. import errno
  16. from common.camera import *
  17. from common.model import *
  18. from common.loss import *
  19. from common.generators import ChunkedGenerator, UnchunkedGenerator
  20. from time import time
  21. from common.utils import deterministic_random
  22. args = parse_args()
  23. print(args)
  24. try:
  25. # Create checkpoint directory if it does not exist
  26. os.makedirs(args.checkpoint)
  27. except OSError as e:
  28. if e.errno != errno.EEXIST:
  29. raise RuntimeError('Unable to create checkpoint directory:', args.checkpoint)
  30. print('Loading dataset...')
  31. dataset_path = 'data/data_3d_' + args.dataset + '.npz'
  32. if args.dataset == 'h36m':
  33. from common.h36m_dataset import Human36mDataset
  34. dataset = Human36mDataset(dataset_path)
  35. elif args.dataset.startswith('humaneva'):
  36. from common.humaneva_dataset import HumanEvaDataset
  37. dataset = HumanEvaDataset(dataset_path)
  38. elif args.dataset.startswith('custom'):
  39. from common.custom_dataset import CustomDataset
  40. dataset = CustomDataset('data/data_2d_' + args.dataset + '_' + args.keypoints + '.npz')
  41. else:
  42. raise KeyError('Invalid dataset')
  43. print('Preparing data...')
  44. for subject in dataset.subjects():
  45. for action in dataset[subject].keys():
  46. anim = dataset[subject][action]
  47. if 'positions' in anim:
  48. positions_3d = []
  49. for cam in anim['cameras']:
  50. pos_3d = world_to_camera(anim['positions'], R=cam['orientation'], t=cam['translation'])
  51. pos_3d[:, 1:] -= pos_3d[:, :1] # Remove global offset, but keep trajectory in first position
  52. positions_3d.append(pos_3d)
  53. anim['positions_3d'] = positions_3d
  54. print('Loading 2D detections...')
  55. keypoints = np.load('data/data_2d_' + args.dataset + '_' + args.keypoints + '.npz', allow_pickle=True)
  56. keypoints_metadata = keypoints['metadata'].item()
  57. keypoints_symmetry = keypoints_metadata['keypoints_symmetry']
  58. kps_left, kps_right = list(keypoints_symmetry[0]), list(keypoints_symmetry[1])
  59. joints_left, joints_right = list(dataset.skeleton().joints_left()), list(dataset.skeleton().joints_right())
  60. keypoints = keypoints['positions_2d'].item()
  61. for subject in dataset.subjects():
  62. assert subject in keypoints, 'Subject {} is missing from the 2D detections dataset'.format(subject)
  63. for action in dataset[subject].keys():
  64. assert action in keypoints[subject], 'Action {} of subject {} is missing from the 2D detections dataset'.format(action, subject)
  65. if 'positions_3d' not in dataset[subject][action]:
  66. continue
  67. for cam_idx in range(len(keypoints[subject][action])):
  68. # We check for >= instead of == because some videos in H3.6M contain extra frames
  69. mocap_length = dataset[subject][action]['positions_3d'][cam_idx].shape[0]
  70. assert keypoints[subject][action][cam_idx].shape[0] >= mocap_length
  71. if keypoints[subject][action][cam_idx].shape[0] > mocap_length:
  72. # Shorten sequence
  73. keypoints[subject][action][cam_idx] = keypoints[subject][action][cam_idx][:mocap_length]
  74. assert len(keypoints[subject][action]) == len(dataset[subject][action]['positions_3d'])
  75. for subject in keypoints.keys():
  76. for action in keypoints[subject]:
  77. for cam_idx, kps in enumerate(keypoints[subject][action]):
  78. # Normalize camera frame
  79. cam = dataset.cameras()[subject][cam_idx]
  80. kps[..., :2] = normalize_screen_coordinates(kps[..., :2], w=cam['res_w'], h=cam['res_h'])
  81. keypoints[subject][action][cam_idx] = kps
  82. subjects_train = args.subjects_train.split(',')
  83. subjects_semi = [] if not args.subjects_unlabeled else args.subjects_unlabeled.split(',')
  84. if not args.render:
  85. subjects_test = args.subjects_test.split(',')
  86. else:
  87. subjects_test = [args.viz_subject]
  88. semi_supervised = len(subjects_semi) > 0
  89. if semi_supervised and not dataset.supports_semi_supervised():
  90. raise RuntimeError('Semi-supervised training is not implemented for this dataset')
  91. def fetch(subjects, action_filter=None, subset=1, parse_3d_poses=True):
  92. out_poses_3d = []
  93. out_poses_2d = []
  94. out_camera_params = []
  95. for subject in subjects:
  96. for action in keypoints[subject].keys():
  97. if action_filter is not None:
  98. found = False
  99. for a in action_filter:
  100. if action.startswith(a):
  101. found = True
  102. break
  103. if not found:
  104. continue
  105. poses_2d = keypoints[subject][action]
  106. for i in range(len(poses_2d)): # Iterate across cameras
  107. out_poses_2d.append(poses_2d[i])
  108. if subject in dataset.cameras():
  109. cams = dataset.cameras()[subject]
  110. assert len(cams) == len(poses_2d), 'Camera count mismatch'
  111. for cam in cams:
  112. if 'intrinsic' in cam:
  113. out_camera_params.append(cam['intrinsic'])
  114. if parse_3d_poses and 'positions_3d' in dataset[subject][action]:
  115. poses_3d = dataset[subject][action]['positions_3d']
  116. assert len(poses_3d) == len(poses_2d), 'Camera count mismatch'
  117. for i in range(len(poses_3d)): # Iterate across cameras
  118. out_poses_3d.append(poses_3d[i])
  119. if len(out_camera_params) == 0:
  120. out_camera_params = None
  121. if len(out_poses_3d) == 0:
  122. out_poses_3d = None
  123. stride = args.downsample
  124. if subset < 1:
  125. for i in range(len(out_poses_2d)):
  126. n_frames = int(round(len(out_poses_2d[i])//stride * subset)*stride)
  127. start = deterministic_random(0, len(out_poses_2d[i]) - n_frames + 1, str(len(out_poses_2d[i])))
  128. out_poses_2d[i] = out_poses_2d[i][start:start+n_frames:stride]
  129. if out_poses_3d is not None:
  130. out_poses_3d[i] = out_poses_3d[i][start:start+n_frames:stride]
  131. elif stride > 1:
  132. # Downsample as requested
  133. for i in range(len(out_poses_2d)):
  134. out_poses_2d[i] = out_poses_2d[i][::stride]
  135. if out_poses_3d is not None:
  136. out_poses_3d[i] = out_poses_3d[i][::stride]
  137. return out_camera_params, out_poses_3d, out_poses_2d
  138. action_filter = None if args.actions == '*' else args.actions.split(',')
  139. if action_filter is not None:
  140. print('Selected actions:', action_filter)
  141. cameras_valid, poses_valid, poses_valid_2d = fetch(subjects_test, action_filter)
  142. filter_widths = [int(x) for x in args.architecture.split(',')]
  143. if not args.disable_optimizations and not args.dense and args.stride == 1:
  144. # Use optimized model for single-frame predictions
  145. model_pos_train = TemporalModelOptimized1f(poses_valid_2d[0].shape[-2], poses_valid_2d[0].shape[-1], dataset.skeleton().num_joints(),
  146. filter_widths=filter_widths, causal=args.causal, dropout=args.dropout, channels=args.channels)
  147. else:
  148. # When incompatible settings are detected (stride > 1, dense filters, or disabled optimization) fall back to normal model
  149. model_pos_train = TemporalModel(poses_valid_2d[0].shape[-2], poses_valid_2d[0].shape[-1], dataset.skeleton().num_joints(),
  150. filter_widths=filter_widths, causal=args.causal, dropout=args.dropout, channels=args.channels,
  151. dense=args.dense)
  152. model_pos = TemporalModel(poses_valid_2d[0].shape[-2], poses_valid_2d[0].shape[-1], dataset.skeleton().num_joints(),
  153. filter_widths=filter_widths, causal=args.causal, dropout=args.dropout, channels=args.channels,
  154. dense=args.dense)
  155. receptive_field = model_pos.receptive_field()
  156. print('INFO: Receptive field: {} frames'.format(receptive_field))
  157. pad = (receptive_field - 1) // 2 # Padding on each side
  158. if args.causal:
  159. print('INFO: Using causal convolutions')
  160. causal_shift = pad
  161. else:
  162. causal_shift = 0
  163. model_params = 0
  164. for parameter in model_pos.parameters():
  165. model_params += parameter.numel()
  166. print('INFO: Trainable parameter count:', model_params)
  167. if torch.cuda.is_available():
  168. model_pos = model_pos.cuda()
  169. model_pos_train = model_pos_train.cuda()
  170. if args.resume or args.evaluate:
  171. chk_filename = os.path.join(args.checkpoint, args.resume if args.resume else args.evaluate)
  172. print('Loading checkpoint', chk_filename)
  173. checkpoint = torch.load(chk_filename, map_location=lambda storage, loc: storage)
  174. print('This model was trained for {} epochs'.format(checkpoint['epoch']))
  175. model_pos_train.load_state_dict(checkpoint['model_pos'])
  176. model_pos.load_state_dict(checkpoint['model_pos'])
  177. if args.evaluate and 'model_traj' in checkpoint:
  178. # Load trajectory model if it contained in the checkpoint (e.g. for inference in the wild)
  179. model_traj = TemporalModel(poses_valid_2d[0].shape[-2], poses_valid_2d[0].shape[-1], 1,
  180. filter_widths=filter_widths, causal=args.causal, dropout=args.dropout, channels=args.channels,
  181. dense=args.dense)
  182. if torch.cuda.is_available():
  183. model_traj = model_traj.cuda()
  184. model_traj.load_state_dict(checkpoint['model_traj'])
  185. else:
  186. model_traj = None
  187. test_generator = UnchunkedGenerator(cameras_valid, poses_valid, poses_valid_2d,
  188. pad=pad, causal_shift=causal_shift, augment=False,
  189. kps_left=kps_left, kps_right=kps_right, joints_left=joints_left, joints_right=joints_right)
  190. print('INFO: Testing on {} frames'.format(test_generator.num_frames()))
  191. if not args.evaluate:
  192. cameras_train, poses_train, poses_train_2d = fetch(subjects_train, action_filter, subset=args.subset)
  193. lr = args.learning_rate
  194. if semi_supervised:
  195. cameras_semi, _, poses_semi_2d = fetch(subjects_semi, action_filter, parse_3d_poses=False)
  196. if not args.disable_optimizations and not args.dense and args.stride == 1:
  197. # Use optimized model for single-frame predictions
  198. model_traj_train = TemporalModelOptimized1f(poses_valid_2d[0].shape[-2], poses_valid_2d[0].shape[-1], 1,
  199. filter_widths=filter_widths, causal=args.causal, dropout=args.dropout, channels=args.channels)
  200. else:
  201. # When incompatible settings are detected (stride > 1, dense filters, or disabled optimization) fall back to normal model
  202. model_traj_train = TemporalModel(poses_valid_2d[0].shape[-2], poses_valid_2d[0].shape[-1], 1,
  203. filter_widths=filter_widths, causal=args.causal, dropout=args.dropout, channels=args.channels,
  204. dense=args.dense)
  205. model_traj = TemporalModel(poses_valid_2d[0].shape[-2], poses_valid_2d[0].shape[-1], 1,
  206. filter_widths=filter_widths, causal=args.causal, dropout=args.dropout, channels=args.channels,
  207. dense=args.dense)
  208. if torch.cuda.is_available():
  209. model_traj = model_traj.cuda()
  210. model_traj_train = model_traj_train.cuda()
  211. optimizer = optim.Adam(list(model_pos_train.parameters()) + list(model_traj_train.parameters()),
  212. lr=lr, amsgrad=True)
  213. losses_2d_train_unlabeled = []
  214. losses_2d_train_labeled_eval = []
  215. losses_2d_train_unlabeled_eval = []
  216. losses_2d_valid = []
  217. losses_traj_train = []
  218. losses_traj_train_eval = []
  219. losses_traj_valid = []
  220. else:
  221. optimizer = optim.Adam(model_pos_train.parameters(), lr=lr, amsgrad=True)
  222. lr_decay = args.lr_decay
  223. losses_3d_train = []
  224. losses_3d_train_eval = []
  225. losses_3d_valid = []
  226. epoch = 0
  227. initial_momentum = 0.1
  228. final_momentum = 0.001
  229. train_generator = ChunkedGenerator(args.batch_size//args.stride, cameras_train, poses_train, poses_train_2d, args.stride,
  230. pad=pad, causal_shift=causal_shift, shuffle=True, augment=args.data_augmentation,
  231. kps_left=kps_left, kps_right=kps_right, joints_left=joints_left, joints_right=joints_right)
  232. train_generator_eval = UnchunkedGenerator(cameras_train, poses_train, poses_train_2d,
  233. pad=pad, causal_shift=causal_shift, augment=False)
  234. print('INFO: Training on {} frames'.format(train_generator_eval.num_frames()))
  235. if semi_supervised:
  236. semi_generator = ChunkedGenerator(args.batch_size//args.stride, cameras_semi, None, poses_semi_2d, args.stride,
  237. pad=pad, causal_shift=causal_shift, shuffle=True,
  238. random_seed=4321, augment=args.data_augmentation,
  239. kps_left=kps_left, kps_right=kps_right, joints_left=joints_left, joints_right=joints_right,
  240. endless=True)
  241. semi_generator_eval = UnchunkedGenerator(cameras_semi, None, poses_semi_2d,
  242. pad=pad, causal_shift=causal_shift, augment=False)
  243. print('INFO: Semi-supervision on {} frames'.format(semi_generator_eval.num_frames()))
  244. if args.resume:
  245. epoch = checkpoint['epoch']
  246. if 'optimizer' in checkpoint and checkpoint['optimizer'] is not None:
  247. optimizer.load_state_dict(checkpoint['optimizer'])
  248. train_generator.set_random_state(checkpoint['random_state'])
  249. else:
  250. print('WARNING: this checkpoint does not contain an optimizer state. The optimizer will be reinitialized.')
  251. lr = checkpoint['lr']
  252. if semi_supervised:
  253. model_traj_train.load_state_dict(checkpoint['model_traj'])
  254. model_traj.load_state_dict(checkpoint['model_traj'])
  255. semi_generator.set_random_state(checkpoint['random_state_semi'])
  256. print('** Note: reported losses are averaged over all frames and test-time augmentation is not used here.')
  257. print('** The final evaluation will be carried out after the last training epoch.')
  258. # Pos model only
  259. while epoch < args.epochs:
  260. start_time = time()
  261. epoch_loss_3d_train = 0
  262. epoch_loss_traj_train = 0
  263. epoch_loss_2d_train_unlabeled = 0
  264. N = 0
  265. N_semi = 0
  266. model_pos_train.train()
  267. if semi_supervised:
  268. # Semi-supervised scenario
  269. model_traj_train.train()
  270. for (_, batch_3d, batch_2d), (cam_semi, _, batch_2d_semi) in \
  271. zip(train_generator.next_epoch(), semi_generator.next_epoch()):
  272. # Fall back to supervised training for the first epoch (to avoid instability)
  273. skip = epoch < args.warmup
  274. cam_semi = torch.from_numpy(cam_semi.astype('float32'))
  275. inputs_3d = torch.from_numpy(batch_3d.astype('float32'))
  276. if torch.cuda.is_available():
  277. cam_semi = cam_semi.cuda()
  278. inputs_3d = inputs_3d.cuda()
  279. inputs_traj = inputs_3d[:, :, :1].clone()
  280. inputs_3d[:, :, 0] = 0
  281. # Split point between labeled and unlabeled samples in the batch
  282. split_idx = inputs_3d.shape[0]
  283. inputs_2d = torch.from_numpy(batch_2d.astype('float32'))
  284. inputs_2d_semi = torch.from_numpy(batch_2d_semi.astype('float32'))
  285. if torch.cuda.is_available():
  286. inputs_2d = inputs_2d.cuda()
  287. inputs_2d_semi = inputs_2d_semi.cuda()
  288. inputs_2d_cat = torch.cat((inputs_2d, inputs_2d_semi), dim=0) if not skip else inputs_2d
  289. optimizer.zero_grad()
  290. # Compute 3D poses
  291. predicted_3d_pos_cat = model_pos_train(inputs_2d_cat)
  292. loss_3d_pos = mpjpe(predicted_3d_pos_cat[:split_idx], inputs_3d)
  293. epoch_loss_3d_train += inputs_3d.shape[0]*inputs_3d.shape[1] * loss_3d_pos.item()
  294. N += inputs_3d.shape[0]*inputs_3d.shape[1]
  295. loss_total = loss_3d_pos
  296. # Compute global trajectory
  297. predicted_traj_cat = model_traj_train(inputs_2d_cat)
  298. w = 1 / inputs_traj[:, :, :, 2] # Weight inversely proportional to depth
  299. loss_traj = weighted_mpjpe(predicted_traj_cat[:split_idx], inputs_traj, w)
  300. epoch_loss_traj_train += inputs_3d.shape[0]*inputs_3d.shape[1] * loss_traj.item()
  301. assert inputs_traj.shape[0]*inputs_traj.shape[1] == inputs_3d.shape[0]*inputs_3d.shape[1]
  302. loss_total += loss_traj
  303. if not skip:
  304. # Semi-supervised loss for unlabeled samples
  305. predicted_semi = predicted_3d_pos_cat[split_idx:]
  306. if pad > 0:
  307. target_semi = inputs_2d_semi[:, pad:-pad, :, :2].contiguous()
  308. else:
  309. target_semi = inputs_2d_semi[:, :, :, :2].contiguous()
  310. projection_func = project_to_2d_linear if args.linear_projection else project_to_2d
  311. reconstruction_semi = projection_func(predicted_semi + predicted_traj_cat[split_idx:], cam_semi)
  312. loss_reconstruction = mpjpe(reconstruction_semi, target_semi) # On 2D poses
  313. epoch_loss_2d_train_unlabeled += predicted_semi.shape[0]*predicted_semi.shape[1] * loss_reconstruction.item()
  314. if not args.no_proj:
  315. loss_total += loss_reconstruction
  316. # Bone length term to enforce kinematic constraints
  317. if args.bone_length_term:
  318. dists = predicted_3d_pos_cat[:, :, 1:] - predicted_3d_pos_cat[:, :, dataset.skeleton().parents()[1:]]
  319. bone_lengths = torch.mean(torch.norm(dists, dim=3), dim=1)
  320. penalty = torch.mean(torch.abs(torch.mean(bone_lengths[:split_idx], dim=0) \
  321. - torch.mean(bone_lengths[split_idx:], dim=0)))
  322. loss_total += penalty
  323. N_semi += predicted_semi.shape[0]*predicted_semi.shape[1]
  324. else:
  325. N_semi += 1 # To avoid division by zero
  326. loss_total.backward()
  327. optimizer.step()
  328. losses_traj_train.append(epoch_loss_traj_train / N)
  329. losses_2d_train_unlabeled.append(epoch_loss_2d_train_unlabeled / N_semi)
  330. else:
  331. # Regular supervised scenario
  332. for _, batch_3d, batch_2d in train_generator.next_epoch():
  333. inputs_3d = torch.from_numpy(batch_3d.astype('float32'))
  334. inputs_2d = torch.from_numpy(batch_2d.astype('float32'))
  335. if torch.cuda.is_available():
  336. inputs_3d = inputs_3d.cuda()
  337. inputs_2d = inputs_2d.cuda()
  338. inputs_3d[:, :, 0] = 0
  339. optimizer.zero_grad()
  340. # Predict 3D poses
  341. predicted_3d_pos = model_pos_train(inputs_2d)
  342. loss_3d_pos = mpjpe(predicted_3d_pos, inputs_3d)
  343. epoch_loss_3d_train += inputs_3d.shape[0]*inputs_3d.shape[1] * loss_3d_pos.item()
  344. N += inputs_3d.shape[0]*inputs_3d.shape[1]
  345. loss_total = loss_3d_pos
  346. loss_total.backward()
  347. optimizer.step()
  348. losses_3d_train.append(epoch_loss_3d_train / N)
  349. # End-of-epoch evaluation
  350. with torch.no_grad():
  351. model_pos.load_state_dict(model_pos_train.state_dict())
  352. model_pos.eval()
  353. if semi_supervised:
  354. model_traj.load_state_dict(model_traj_train.state_dict())
  355. model_traj.eval()
  356. epoch_loss_3d_valid = 0
  357. epoch_loss_traj_valid = 0
  358. epoch_loss_2d_valid = 0
  359. N = 0
  360. if not args.no_eval:
  361. # Evaluate on test set
  362. for cam, batch, batch_2d in test_generator.next_epoch():
  363. inputs_3d = torch.from_numpy(batch.astype('float32'))
  364. inputs_2d = torch.from_numpy(batch_2d.astype('float32'))
  365. if torch.cuda.is_available():
  366. inputs_3d = inputs_3d.cuda()
  367. inputs_2d = inputs_2d.cuda()
  368. inputs_traj = inputs_3d[:, :, :1].clone()
  369. inputs_3d[:, :, 0] = 0
  370. # Predict 3D poses
  371. predicted_3d_pos = model_pos(inputs_2d)
  372. loss_3d_pos = mpjpe(predicted_3d_pos, inputs_3d)
  373. epoch_loss_3d_valid += inputs_3d.shape[0]*inputs_3d.shape[1] * loss_3d_pos.item()
  374. N += inputs_3d.shape[0]*inputs_3d.shape[1]
  375. if semi_supervised:
  376. cam = torch.from_numpy(cam.astype('float32'))
  377. if torch.cuda.is_available():
  378. cam = cam.cuda()
  379. predicted_traj = model_traj(inputs_2d)
  380. loss_traj = mpjpe(predicted_traj, inputs_traj)
  381. epoch_loss_traj_valid += inputs_traj.shape[0]*inputs_traj.shape[1] * loss_traj.item()
  382. assert inputs_traj.shape[0]*inputs_traj.shape[1] == inputs_3d.shape[0]*inputs_3d.shape[1]
  383. if pad > 0:
  384. target = inputs_2d[:, pad:-pad, :, :2].contiguous()
  385. else:
  386. target = inputs_2d[:, :, :, :2].contiguous()
  387. reconstruction = project_to_2d(predicted_3d_pos + predicted_traj, cam)
  388. loss_reconstruction = mpjpe(reconstruction, target) # On 2D poses
  389. epoch_loss_2d_valid += reconstruction.shape[0]*reconstruction.shape[1] * loss_reconstruction.item()
  390. assert reconstruction.shape[0]*reconstruction.shape[1] == inputs_3d.shape[0]*inputs_3d.shape[1]
  391. losses_3d_valid.append(epoch_loss_3d_valid / N)
  392. if semi_supervised:
  393. losses_traj_valid.append(epoch_loss_traj_valid / N)
  394. losses_2d_valid.append(epoch_loss_2d_valid / N)
  395. # Evaluate on training set, this time in evaluation mode
  396. epoch_loss_3d_train_eval = 0
  397. epoch_loss_traj_train_eval = 0
  398. epoch_loss_2d_train_labeled_eval = 0
  399. N = 0
  400. for cam, batch, batch_2d in train_generator_eval.next_epoch():
  401. if batch_2d.shape[1] == 0:
  402. # This can only happen when downsampling the dataset
  403. continue
  404. inputs_3d = torch.from_numpy(batch.astype('float32'))
  405. inputs_2d = torch.from_numpy(batch_2d.astype('float32'))
  406. if torch.cuda.is_available():
  407. inputs_3d = inputs_3d.cuda()
  408. inputs_2d = inputs_2d.cuda()
  409. inputs_traj = inputs_3d[:, :, :1].clone()
  410. inputs_3d[:, :, 0] = 0
  411. # Compute 3D poses
  412. predicted_3d_pos = model_pos(inputs_2d)
  413. loss_3d_pos = mpjpe(predicted_3d_pos, inputs_3d)
  414. epoch_loss_3d_train_eval += inputs_3d.shape[0]*inputs_3d.shape[1] * loss_3d_pos.item()
  415. N += inputs_3d.shape[0]*inputs_3d.shape[1]
  416. if semi_supervised:
  417. cam = torch.from_numpy(cam.astype('float32'))
  418. if torch.cuda.is_available():
  419. cam = cam.cuda()
  420. predicted_traj = model_traj(inputs_2d)
  421. loss_traj = mpjpe(predicted_traj, inputs_traj)
  422. epoch_loss_traj_train_eval += inputs_traj.shape[0]*inputs_traj.shape[1] * loss_traj.item()
  423. assert inputs_traj.shape[0]*inputs_traj.shape[1] == inputs_3d.shape[0]*inputs_3d.shape[1]
  424. if pad > 0:
  425. target = inputs_2d[:, pad:-pad, :, :2].contiguous()
  426. else:
  427. target = inputs_2d[:, :, :, :2].contiguous()
  428. reconstruction = project_to_2d(predicted_3d_pos + predicted_traj, cam)
  429. loss_reconstruction = mpjpe(reconstruction, target)
  430. epoch_loss_2d_train_labeled_eval += reconstruction.shape[0]*reconstruction.shape[1] * loss_reconstruction.item()
  431. assert reconstruction.shape[0]*reconstruction.shape[1] == inputs_3d.shape[0]*inputs_3d.shape[1]
  432. losses_3d_train_eval.append(epoch_loss_3d_train_eval / N)
  433. if semi_supervised:
  434. losses_traj_train_eval.append(epoch_loss_traj_train_eval / N)
  435. losses_2d_train_labeled_eval.append(epoch_loss_2d_train_labeled_eval / N)
  436. # Evaluate 2D loss on unlabeled training set (in evaluation mode)
  437. epoch_loss_2d_train_unlabeled_eval = 0
  438. N_semi = 0
  439. if semi_supervised:
  440. for cam, _, batch_2d in semi_generator_eval.next_epoch():
  441. cam = torch.from_numpy(cam.astype('float32'))
  442. inputs_2d_semi = torch.from_numpy(batch_2d.astype('float32'))
  443. if torch.cuda.is_available():
  444. cam = cam.cuda()
  445. inputs_2d_semi = inputs_2d_semi.cuda()
  446. predicted_3d_pos_semi = model_pos(inputs_2d_semi)
  447. predicted_traj_semi = model_traj(inputs_2d_semi)
  448. if pad > 0:
  449. target_semi = inputs_2d_semi[:, pad:-pad, :, :2].contiguous()
  450. else:
  451. target_semi = inputs_2d_semi[:, :, :, :2].contiguous()
  452. reconstruction_semi = project_to_2d(predicted_3d_pos_semi + predicted_traj_semi, cam)
  453. loss_reconstruction_semi = mpjpe(reconstruction_semi, target_semi)
  454. epoch_loss_2d_train_unlabeled_eval += reconstruction_semi.shape[0]*reconstruction_semi.shape[1] \
  455. * loss_reconstruction_semi.item()
  456. N_semi += reconstruction_semi.shape[0]*reconstruction_semi.shape[1]
  457. losses_2d_train_unlabeled_eval.append(epoch_loss_2d_train_unlabeled_eval / N_semi)
  458. elapsed = (time() - start_time)/60
  459. if args.no_eval:
  460. print('[%d] time %.2f lr %f 3d_train %f' % (
  461. epoch + 1,
  462. elapsed,
  463. lr,
  464. losses_3d_train[-1] * 1000))
  465. else:
  466. if semi_supervised:
  467. print('[%d] time %.2f lr %f 3d_train %f 3d_eval %f traj_eval %f 3d_valid %f '
  468. 'traj_valid %f 2d_train_sup %f 2d_train_unsup %f 2d_valid %f' % (
  469. epoch + 1,
  470. elapsed,
  471. lr,
  472. losses_3d_train[-1] * 1000,
  473. losses_3d_train_eval[-1] * 1000,
  474. losses_traj_train_eval[-1] * 1000,
  475. losses_3d_valid[-1] * 1000,
  476. losses_traj_valid[-1] * 1000,
  477. losses_2d_train_labeled_eval[-1],
  478. losses_2d_train_unlabeled_eval[-1],
  479. losses_2d_valid[-1]))
  480. else:
  481. print('[%d] time %.2f lr %f 3d_train %f 3d_eval %f 3d_valid %f' % (
  482. epoch + 1,
  483. elapsed,
  484. lr,
  485. losses_3d_train[-1] * 1000,
  486. losses_3d_train_eval[-1] * 1000,
  487. losses_3d_valid[-1] *1000))
  488. # Decay learning rate exponentially
  489. lr *= lr_decay
  490. for param_group in optimizer.param_groups:
  491. param_group['lr'] *= lr_decay
  492. epoch += 1
  493. # Decay BatchNorm momentum
  494. momentum = initial_momentum * np.exp(-epoch/args.epochs * np.log(initial_momentum/final_momentum))
  495. model_pos_train.set_bn_momentum(momentum)
  496. if semi_supervised:
  497. model_traj_train.set_bn_momentum(momentum)
  498. # Save checkpoint if necessary
  499. if epoch % args.checkpoint_frequency == 0:
  500. chk_path = os.path.join(args.checkpoint, 'epoch_{}.bin'.format(epoch))
  501. print('Saving checkpoint to', chk_path)
  502. torch.save({
  503. 'epoch': epoch,
  504. 'lr': lr,
  505. 'random_state': train_generator.random_state(),
  506. 'optimizer': optimizer.state_dict(),
  507. 'model_pos': model_pos_train.state_dict(),
  508. 'model_traj': model_traj_train.state_dict() if semi_supervised else None,
  509. 'random_state_semi': semi_generator.random_state() if semi_supervised else None,
  510. }, chk_path)
  511. # Save training curves after every epoch, as .png images (if requested)
  512. if args.export_training_curves and epoch > 3:
  513. if 'matplotlib' not in sys.modules:
  514. import matplotlib
  515. matplotlib.use('Agg')
  516. import matplotlib.pyplot as plt
  517. plt.figure()
  518. epoch_x = np.arange(3, len(losses_3d_train)) + 1
  519. plt.plot(epoch_x, losses_3d_train[3:], '--', color='C0')
  520. plt.plot(epoch_x, losses_3d_train_eval[3:], color='C0')
  521. plt.plot(epoch_x, losses_3d_valid[3:], color='C1')
  522. plt.legend(['3d train', '3d train (eval)', '3d valid (eval)'])
  523. plt.ylabel('MPJPE (m)')
  524. plt.xlabel('Epoch')
  525. plt.xlim((3, epoch))
  526. plt.savefig(os.path.join(args.checkpoint, 'loss_3d.png'))
  527. if semi_supervised:
  528. plt.figure()
  529. plt.plot(epoch_x, losses_traj_train[3:], '--', color='C0')
  530. plt.plot(epoch_x, losses_traj_train_eval[3:], color='C0')
  531. plt.plot(epoch_x, losses_traj_valid[3:], color='C1')
  532. plt.legend(['traj. train', 'traj. train (eval)', 'traj. valid (eval)'])
  533. plt.ylabel('Mean distance (m)')
  534. plt.xlabel('Epoch')
  535. plt.xlim((3, epoch))
  536. plt.savefig(os.path.join(args.checkpoint, 'loss_traj.png'))
  537. plt.figure()
  538. plt.plot(epoch_x, losses_2d_train_labeled_eval[3:], color='C0')
  539. plt.plot(epoch_x, losses_2d_train_unlabeled[3:], '--', color='C1')
  540. plt.plot(epoch_x, losses_2d_train_unlabeled_eval[3:], color='C1')
  541. plt.plot(epoch_x, losses_2d_valid[3:], color='C2')
  542. plt.legend(['2d train labeled (eval)', '2d train unlabeled', '2d train unlabeled (eval)', '2d valid (eval)'])
  543. plt.ylabel('MPJPE (2D)')
  544. plt.xlabel('Epoch')
  545. plt.xlim((3, epoch))
  546. plt.savefig(os.path.join(args.checkpoint, 'loss_2d.png'))
  547. plt.close('all')
  548. # Evaluate
  549. def evaluate(test_generator, action=None, return_predictions=False, use_trajectory_model=False):
  550. epoch_loss_3d_pos = 0
  551. epoch_loss_3d_pos_procrustes = 0
  552. epoch_loss_3d_pos_scale = 0
  553. epoch_loss_3d_vel = 0
  554. with torch.no_grad():
  555. if not use_trajectory_model:
  556. model_pos.eval()
  557. else:
  558. model_traj.eval()
  559. N = 0
  560. for _, batch, batch_2d in test_generator.next_epoch():
  561. inputs_2d = torch.from_numpy(batch_2d.astype('float32'))
  562. if torch.cuda.is_available():
  563. inputs_2d = inputs_2d.cuda()
  564. # Positional model
  565. if not use_trajectory_model:
  566. predicted_3d_pos = model_pos(inputs_2d)
  567. else:
  568. predicted_3d_pos = model_traj(inputs_2d)
  569. # Test-time augmentation (if enabled)
  570. if test_generator.augment_enabled():
  571. # Undo flipping and take average with non-flipped version
  572. predicted_3d_pos[1, :, :, 0] *= -1
  573. if not use_trajectory_model:
  574. predicted_3d_pos[1, :, joints_left + joints_right] = predicted_3d_pos[1, :, joints_right + joints_left]
  575. predicted_3d_pos = torch.mean(predicted_3d_pos, dim=0, keepdim=True)
  576. if return_predictions:
  577. return predicted_3d_pos.squeeze(0).cpu().numpy()
  578. inputs_3d = torch.from_numpy(batch.astype('float32'))
  579. if torch.cuda.is_available():
  580. inputs_3d = inputs_3d.cuda()
  581. inputs_3d[:, :, 0] = 0
  582. if test_generator.augment_enabled():
  583. inputs_3d = inputs_3d[:1]
  584. error = mpjpe(predicted_3d_pos, inputs_3d)
  585. epoch_loss_3d_pos_scale += inputs_3d.shape[0]*inputs_3d.shape[1] * n_mpjpe(predicted_3d_pos, inputs_3d).item()
  586. epoch_loss_3d_pos += inputs_3d.shape[0]*inputs_3d.shape[1] * error.item()
  587. N += inputs_3d.shape[0] * inputs_3d.shape[1]
  588. inputs = inputs_3d.cpu().numpy().reshape(-1, inputs_3d.shape[-2], inputs_3d.shape[-1])
  589. predicted_3d_pos = predicted_3d_pos.cpu().numpy().reshape(-1, inputs_3d.shape[-2], inputs_3d.shape[-1])
  590. epoch_loss_3d_pos_procrustes += inputs_3d.shape[0]*inputs_3d.shape[1] * p_mpjpe(predicted_3d_pos, inputs)
  591. # Compute velocity error
  592. epoch_loss_3d_vel += inputs_3d.shape[0]*inputs_3d.shape[1] * mean_velocity_error(predicted_3d_pos, inputs)
  593. if action is None:
  594. print('----------')
  595. else:
  596. print('----'+action+'----')
  597. e1 = (epoch_loss_3d_pos / N)*1000
  598. e2 = (epoch_loss_3d_pos_procrustes / N)*1000
  599. e3 = (epoch_loss_3d_pos_scale / N)*1000
  600. ev = (epoch_loss_3d_vel / N)*1000
  601. print('Test time augmentation:', test_generator.augment_enabled())
  602. print('Protocol #1 Error (MPJPE):', e1, 'mm')
  603. print('Protocol #2 Error (P-MPJPE):', e2, 'mm')
  604. print('Protocol #3 Error (N-MPJPE):', e3, 'mm')
  605. print('Velocity Error (MPJVE):', ev, 'mm')
  606. print('----------')
  607. return e1, e2, e3, ev
  608. if args.render:
  609. print('Rendering...')
  610. input_keypoints = keypoints[args.viz_subject][args.viz_action][args.viz_camera].copy()
  611. ground_truth = None
  612. if args.viz_subject in dataset.subjects() and args.viz_action in dataset[args.viz_subject]:
  613. if 'positions_3d' in dataset[args.viz_subject][args.viz_action]:
  614. ground_truth = dataset[args.viz_subject][args.viz_action]['positions_3d'][args.viz_camera].copy()
  615. if ground_truth is None:
  616. print('INFO: this action is unlabeled. Ground truth will not be rendered.')
  617. gen = UnchunkedGenerator(None, None, [input_keypoints],
  618. pad=pad, causal_shift=causal_shift, augment=args.test_time_augmentation,
  619. kps_left=kps_left, kps_right=kps_right, joints_left=joints_left, joints_right=joints_right)
  620. prediction = evaluate(gen, return_predictions=True)
  621. if model_traj is not None and ground_truth is None:
  622. prediction_traj = evaluate(gen, return_predictions=True, use_trajectory_model=True)
  623. prediction += prediction_traj
  624. if args.viz_export is not None:
  625. print('Exporting joint positions to', args.viz_export)
  626. # Predictions are in camera space
  627. np.save(args.viz_export, prediction)
  628. if args.viz_output is not None:
  629. if ground_truth is not None:
  630. # Reapply trajectory
  631. trajectory = ground_truth[:, :1]
  632. ground_truth[:, 1:] += trajectory
  633. prediction += trajectory
  634. # Invert camera transformation
  635. cam = dataset.cameras()[args.viz_subject][args.viz_camera]
  636. if ground_truth is not None:
  637. prediction = camera_to_world(prediction, R=cam['orientation'], t=cam['translation'])
  638. ground_truth = camera_to_world(ground_truth, R=cam['orientation'], t=cam['translation'])
  639. else:
  640. # If the ground truth is not available, take the camera extrinsic params from a random subject.
  641. # They are almost the same, and anyway, we only need this for visualization purposes.
  642. for subject in dataset.cameras():
  643. if 'orientation' in dataset.cameras()[subject][args.viz_camera]:
  644. rot = dataset.cameras()[subject][args.viz_camera]['orientation']
  645. break
  646. prediction = camera_to_world(prediction, R=rot, t=0)
  647. # We don't have the trajectory, but at least we can rebase the height
  648. prediction[:, :, 2] -= np.min(prediction[:, :, 2])
  649. anim_output = {'Reconstruction': prediction}
  650. if ground_truth is not None and not args.viz_no_ground_truth:
  651. anim_output['Ground truth'] = ground_truth
  652. input_keypoints = image_coordinates(input_keypoints[..., :2], w=cam['res_w'], h=cam['res_h'])
  653. from common.visualization import render_animation
  654. render_animation(input_keypoints, keypoints_metadata, anim_output,
  655. dataset.skeleton(), dataset.fps(), args.viz_bitrate, cam['azimuth'], args.viz_output,
  656. limit=args.viz_limit, downsample=args.viz_downsample, size=args.viz_size,
  657. input_video_path=args.viz_video, viewport=(cam['res_w'], cam['res_h']),
  658. input_video_skip=args.viz_skip)
  659. else:
  660. print('Evaluating...')
  661. all_actions = {}
  662. all_actions_by_subject = {}
  663. for subject in subjects_test:
  664. if subject not in all_actions_by_subject:
  665. all_actions_by_subject[subject] = {}
  666. for action in dataset[subject].keys():
  667. action_name = action.split(' ')[0]
  668. if action_name not in all_actions:
  669. all_actions[action_name] = []
  670. if action_name not in all_actions_by_subject[subject]:
  671. all_actions_by_subject[subject][action_name] = []
  672. all_actions[action_name].append((subject, action))
  673. all_actions_by_subject[subject][action_name].append((subject, action))
  674. def fetch_actions(actions):
  675. out_poses_3d = []
  676. out_poses_2d = []
  677. for subject, action in actions:
  678. poses_2d = keypoints[subject][action]
  679. for i in range(len(poses_2d)): # Iterate across cameras
  680. out_poses_2d.append(poses_2d[i])
  681. poses_3d = dataset[subject][action]['positions_3d']
  682. assert len(poses_3d) == len(poses_2d), 'Camera count mismatch'
  683. for i in range(len(poses_3d)): # Iterate across cameras
  684. out_poses_3d.append(poses_3d[i])
  685. stride = args.downsample
  686. if stride > 1:
  687. # Downsample as requested
  688. for i in range(len(out_poses_2d)):
  689. out_poses_2d[i] = out_poses_2d[i][::stride]
  690. if out_poses_3d is not None:
  691. out_poses_3d[i] = out_poses_3d[i][::stride]
  692. return out_poses_3d, out_poses_2d
  693. def run_evaluation(actions, action_filter=None):
  694. errors_p1 = []
  695. errors_p2 = []
  696. errors_p3 = []
  697. errors_vel = []
  698. for action_key in actions.keys():
  699. if action_filter is not None:
  700. found = False
  701. for a in action_filter:
  702. if action_key.startswith(a):
  703. found = True
  704. break
  705. if not found:
  706. continue
  707. poses_act, poses_2d_act = fetch_actions(actions[action_key])
  708. gen = UnchunkedGenerator(None, poses_act, poses_2d_act,
  709. pad=pad, causal_shift=causal_shift, augment=args.test_time_augmentation,
  710. kps_left=kps_left, kps_right=kps_right, joints_left=joints_left, joints_right=joints_right)
  711. e1, e2, e3, ev = evaluate(gen, action_key)
  712. errors_p1.append(e1)
  713. errors_p2.append(e2)
  714. errors_p3.append(e3)
  715. errors_vel.append(ev)
  716. print('Protocol #1 (MPJPE) action-wise average:', round(np.mean(errors_p1), 1), 'mm')
  717. print('Protocol #2 (P-MPJPE) action-wise average:', round(np.mean(errors_p2), 1), 'mm')
  718. print('Protocol #3 (N-MPJPE) action-wise average:', round(np.mean(errors_p3), 1), 'mm')
  719. print('Velocity (MPJVE) action-wise average:', round(np.mean(errors_vel), 2), 'mm')
  720. if not args.by_subject:
  721. run_evaluation(all_actions, action_filter)
  722. else:
  723. for subject in all_actions_by_subject.keys():
  724. print('Evaluating on subject', subject)
  725. run_evaluation(all_actions_by_subject[subject], action_filter)
  726. print('')