(self, data_list)
| 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, \ |
| 77 | ori_pixel_ids, pixel_ids, gripper_direction_camera, gripper_forward_direction_camera, True, False) |
| 78 | self.false_data[cur_primact_type].append(cur_data) |
| 79 | |
| 80 | # load neg-direction false data |
| 81 | if not self.no_neg_dir_data: |
| 82 | cur_data = (cur_dir, cur_shape_id, cur_category, cur_cnt_id, cur_trial_id, \ |
| 83 | ori_pixel_ids, pixel_ids, -gripper_direction_camera, gripper_forward_direction_camera, False, False) |
| 84 | self.false_data[cur_primact_type].append(cur_data) |
| 85 | |
| 86 | # delete data if buffer full |
| 87 | if self.buffer_max_num is not None: |
| 88 | for primact_type in self.primact_types: |
| 89 | if len(self.true_data[primact_type]) > self.buffer_max_num: |
| 90 | self.true_data[primact_type] = self.true_data[primact_type][-self.buffer_max_num:] |
| 91 | if len(self.false_data[primact_type]) > self.buffer_max_num: |
| 92 | self.false_data[primact_type] = self.false_data[primact_type][-self.buffer_max_num:] |
| 93 | |
| 94 | def check_success(self, result_data, primact_type): |
| 95 | if result_data['result'] != 'VALID': |
no test coverage detected