run.py 45 KB

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