| 16 | |
| 17 | |
| 18 | class Fusion(data.Dataset): |
| 19 | def __init__(self, opt, dataset, root_path, train=True): |
| 20 | self.data_type = opt.dataset # h36m |
| 21 | self.train = train |
| 22 | self.keypoints_name = opt.keypoints |
| 23 | self.root_path = root_path |
| 24 | |
| 25 | self.train_list = opt.subjects_train.split(",") |
| 26 | self.test_list = opt.subjects_test.split(",") |
| 27 | self.action_filter = None if opt.actions == "*" else opt.actions.split(",") |
| 28 | self.downsample = opt.downsample |
| 29 | self.subset = opt.subset |
| 30 | self.stride = opt.stride |
| 31 | self.crop_uv = opt.crop_uv |
| 32 | self.test_aug = opt.test_augmentation |
| 33 | self.pad = opt.pad # 0 |
| 34 | if self.train: |
| 35 | self.keypoints = self.prepare_data(dataset, self.train_list) |
| 36 | self.cameras_train, self.poses_train, self.poses_train_2d = self.fetch( |
| 37 | dataset, self.train_list, subset=self.subset, views=opt.train_views |
| 38 | ) |
| 39 | |
| 40 | self.generator = ChunkedGenerator( |
| 41 | opt.batch_size // opt.stride, |
| 42 | self.cameras_train, |
| 43 | self.poses_train, |
| 44 | self.poses_train_2d, |
| 45 | self.stride, |
| 46 | pad=self.pad, |
| 47 | augment=opt.data_augmentation, |
| 48 | reverse_aug=opt.reverse_augmentation, |
| 49 | kps_left=self.kps_left, |
| 50 | kps_right=self.kps_right, |
| 51 | joints_left=self.joints_left, |
| 52 | joints_right=self.joints_right, |
| 53 | out_all=opt.out_all, |
| 54 | ) |
| 55 | print("INFO: Training on {} frames".format(self.generator.num_frames())) |
| 56 | else: |
| 57 | self.keypoints = self.prepare_data(dataset, self.test_list) |
| 58 | self.cameras_test, self.poses_test, self.poses_test_2d = self.fetch( |
| 59 | dataset, self.test_list, subset=self.subset, views=opt.test_views |
| 60 | ) |
| 61 | |
| 62 | self.generator = ChunkedGenerator( |
| 63 | opt.batch_size // opt.stride, |
| 64 | self.cameras_test, |
| 65 | self.poses_test, |
| 66 | self.poses_test_2d, |
| 67 | pad=self.pad, |
| 68 | augment=False, |
| 69 | kps_left=self.kps_left, |
| 70 | kps_right=self.kps_right, |
| 71 | joints_left=self.joints_left, |
| 72 | joints_right=self.joints_right, |
| 73 | ) |
| 74 | self.key_index = self.generator.saved_index |
| 75 | print("INFO: Testing on {} frames".format(self.generator.num_frames())) |