| 5 | import logging |
| 6 | |
| 7 | def parse_args(): |
| 8 | parser = argparse.ArgumentParser() |
| 9 | |
| 10 | parser.add_argument('--layers', default=3, type=int) |
| 11 | parser.add_argument('--channel', default=512, type=int) |
| 12 | parser.add_argument('--d_hid', default=1024, type=int) |
| 13 | parser.add_argument('--token_dim', default=256, type=int) |
| 14 | parser.add_argument('--dataset', type=str, default='h36m') |
| 15 | parser.add_argument('--keypoints', default='cpn_ft_h36m_dbb', type=str) |
| 16 | parser.add_argument('--data_augmentation', type=int, default=1) |
| 17 | parser.add_argument('--reverse_augmentation', type=bool, default=False) |
| 18 | parser.add_argument('--test_augmentation', type=bool, default=True) |
| 19 | parser.add_argument('--crop_uv', type=int, default=0) |
| 20 | parser.add_argument('--root_path', type=str, default='dataset/') |
| 21 | parser.add_argument('--actions', default='*', type=str) |
| 22 | parser.add_argument('--downsample', default=1, type=int) |
| 23 | parser.add_argument('--subset', default=1, type=float) |
| 24 | parser.add_argument('--stride', default=1, type=int) |
| 25 | parser.add_argument('--gpu', default='0', type=str) |
| 26 | parser.add_argument('--train', default=1, type=int) |
| 27 | parser.add_argument('--test', action='store_true') |
| 28 | parser.add_argument('--nepoch', type=int, default=20) |
| 29 | parser.add_argument('--batch_size', type=int, default=256) |
| 30 | parser.add_argument('--lr', type=float, default=1e-3) |
| 31 | parser.add_argument('--lr_refine', type=float, default=1e-5) |
| 32 | parser.add_argument('--lr_decay_large', type=float, default=0.5) |
| 33 | parser.add_argument('--lr_decay_epoch', type=int, default=5) |
| 34 | parser.add_argument('--workers', type=int, default=8) |
| 35 | parser.add_argument('-lrd', '--lr_decay', default=0.95, type=float) |
| 36 | parser.add_argument('--frames', type=int, default=1) |
| 37 | parser.add_argument('--pad', type=int, default=0) |
| 38 | parser.add_argument('--refine', action='store_true') |
| 39 | parser.add_argument('--refine_reload', action='store_true') |
| 40 | parser.add_argument('--checkpoint', type=str, default='') |
| 41 | parser.add_argument('--previous_dir', type=str, default='') |
| 42 | parser.add_argument('--n_joints', type=int, default=17) |
| 43 | parser.add_argument('--out_joints', type=int, default=17) |
| 44 | parser.add_argument('--out_all', type=int, default=1) |
| 45 | parser.add_argument('--out_channels', type=int, default=3) |
| 46 | parser.add_argument('--previous_best', type=float, default= math.inf) |
| 47 | parser.add_argument('--previous_name', type=str, default='') |
| 48 | parser.add_argument('--previous_refine_name', type=str, default='') |
| 49 | |
| 50 | args = parser.parse_args() |
| 51 | |
| 52 | if args.test: |
| 53 | args.train = 0 |
| 54 | |
| 55 | args.pad = (args.frames-1) // 2 |
| 56 | |
| 57 | args.root_joint = 0 |
| 58 | if args.dataset == 'h36m': |
| 59 | args.subjects_train = 'S1,S5,S6,S7,S8' |
| 60 | args.subjects_test = 'S9,S11' |
| 61 | |
| 62 | args.n_joints = 17 |
| 63 | args.out_joints = 17 |
| 64 | |