MCPcopy Create free account
hub / github.com/daerduoCarey/where2act / load_data

Method load_data

code/data.py:46–92  ·  view source on GitHub ↗
(self, data_list)

Source from the content-addressed store, hash-verified

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':

Callers 2

trainFunction · 0.95
trainFunction · 0.95

Calls 1

check_successMethod · 0.95

Tested by

no test coverage detected