| 17 | |
| 18 | |
| 19 | class SAPIENVisionDataset(data.Dataset): |
| 20 | |
| 21 | def __init__(self, primact_types, category_types, data_features, buffer_max_num, \ |
| 22 | abs_thres=0.01, rel_thres=0.5, dp_thres=0.5, img_size=224, \ |
| 23 | no_true_false_equal=False, no_neg_dir_data=False, only_true_data=False): |
| 24 | self.primact_types = primact_types |
| 25 | self.category_types = category_types |
| 26 | |
| 27 | self.buffer_max_num = buffer_max_num |
| 28 | self.img_size = img_size |
| 29 | self.abs_thres = abs_thres |
| 30 | self.rel_thres = rel_thres |
| 31 | self.dp_thres = dp_thres |
| 32 | self.no_true_false_equal = no_true_false_equal |
| 33 | self.no_neg_dir_data = no_neg_dir_data |
| 34 | self.only_true_data = only_true_data |
| 35 | |
| 36 | # data buffer |
| 37 | self.true_data = dict() |
| 38 | self.false_data = dict() |
| 39 | for primact_type in primact_types: |
| 40 | self.true_data[primact_type] = [] |
| 41 | self.false_data[primact_type] = [] |
| 42 | |
| 43 | # data features |
| 44 | self.data_features = data_features |
| 45 | |
| 46 | def load_data(self, data_list): |
| 47 | bar = ProgressBar() |
| 48 | for i in bar(range(len(data_list))): |
| 49 | cur_dir = data_list[i] |
| 50 | cur_shape_id, cur_category, cur_cnt_id, cur_primact_type, cur_trial_id = cur_dir.split('/')[-1].split('_') |
| 51 | |
| 52 | if cur_primact_type not in self.primact_types: |
| 53 | continue |
| 54 | |
| 55 | if cur_category not in self.category_types: |
| 56 | continue |
| 57 | |
| 58 | with open(os.path.join(cur_dir, 'result.json'), 'r') as fin: |
| 59 | result_data = json.load(fin) |
| 60 | |
| 61 | gripper_direction_camera = np.array(result_data['gripper_direction_camera'], dtype=np.float32) |
| 62 | gripper_forward_direction_camera = np.array(result_data['gripper_forward_direction_camera'], dtype=np.float32) |
| 63 | |
| 64 | ori_pixel_ids = np.array(result_data['pixel_locs'], dtype=np.int32) |
| 65 | pixel_ids = np.round(np.array(result_data['pixel_locs'], dtype=np.float32) / 448 * self.img_size).astype(np.int32) |
| 66 | |
| 67 | success = self.check_success(result_data, cur_primact_type) |
| 68 | |
| 69 | # load original data |
| 70 | if success: |
| 71 | cur_data = (cur_dir, cur_shape_id, cur_category, cur_cnt_id, cur_trial_id, \ |
| 72 | ori_pixel_ids, pixel_ids, gripper_direction_camera, gripper_forward_direction_camera, True, True) |
| 73 | self.true_data[cur_primact_type].append(cur_data) |
| 74 | else: |
| 75 | if not self.only_true_data: |
| 76 | cur_data = (cur_dir, cur_shape_id, cur_category, cur_cnt_id, cur_trial_id, \ |