MCPcopy Create free account
hub / github.com/Relento/lego_release / evaluate_set

Function evaluate_set

eval.py:122–298  ·  view source on GitHub ↗
(opt, model, set_path, meters, is_subm=False, subm_cbricks=None, subm_meters=None)

Source from the content-addressed store, hash-verified

120
121
122def evaluate_set(opt, model, set_path, meters, is_subm=False, subm_cbricks=None, subm_meters=None):
123 result_list = []
124 opt.load_set = set_path
125 opt.serial_batches = True
126 autoregressive = opt.autoregressive_inference
127 if autoregressive:
128 opt.batch_size = 1
129 current_bricks_pc = BricksPC(grid_size=(65, 65, 65), record_parents=False)
130 if subm_cbricks is not None:
131 subm_ct = 0
132
133 meters_this = Meters()
134 dataset = create_dataset(opt) # create a dataset given opt.dataset_mode and other options
135 print(f'Loading from {set_path}, number of steps {len(dataset)}')
136
137 if is_subm:
138 oracle_ct = 0 # disable oracle in subm
139 else:
140 oracle_ct = len(dataset) * opt.oracle_percentage
141 for i, data in tqdm(enumerate(dataset), total=len(dataset) // opt.batch_size):
142 if i > oracle_ct and autoregressive:
143 data[0]['obj_occ_prev'] = torch.as_tensor(current_bricks_pc.get_occ_with_rotation()[0])[None]
144
145 data[0]['bricks'] = [copy.deepcopy(current_bricks_pc)]
146 transforms = pt.Transform3d().translate(-data[0]['obj_center']).scale(data[0]['obj_scale'][0]).cuda()
147 cameras = get_cameras(azim=data[0]['azim'][0], elev=data[0]['elev'][0])
148 data[1][0]['conns'] = recompute_conns(
149 current_bricks_pc, op_type=0, transforms=transforms, cameras=cameras, scale=get_scale())
150
151 if subm_cbricks is not None:
152 # Use the predicted submodules as input
153 for j, bid_counter in enumerate((data[0]['bid_counter'])):
154 has_subm = False
155 for k, (bid, ct) in enumerate(bid_counter):
156 if bid < 0:
157 has_subm = True
158 data[0]['brick_occs'][j][k] = get_brick_occ(opt, subm_cbricks[subm_ct])
159 data[0]['cbrick'][j][-bid - 1] = subm_cbricks[subm_ct]
160
161 if has_subm:
162 subm_ct += 1
163
164 model.set_input(data)
165 model.test()
166 losses = model.get_current_losses()
167 data, targets = data
168 for name, v in losses.items():
169 obj_sum = sum(data['reg_mask'][i].sum().item() for i in range(len(data['reg_mask'])))
170 meters.update(name, v * obj_sum, obj_sum)
171
172 def zip_keys(exs, features=['trans', 'rot_decoded', 'bid', 'bid_decoded'], n=-1):
173 res = []
174 if n == -1:
175 n = len(exs[features[0]])
176 for idx in range(n):
177 ex = {}
178 for k in features:
179 if k in exs:

Callers 1

eval.pyFile · 0.85

Calls 15

get_occ_with_rotationMethod · 0.95
updateMethod · 0.95
add_brickMethod · 0.95
BricksPCClass · 0.90
MetersClass · 0.90
create_datasetFunction · 0.90
get_camerasFunction · 0.90
recompute_connsFunction · 0.90
get_scaleFunction · 0.90
get_brick_classFunction · 0.90
get_cbrick_keypointFunction · 0.90
add_cbrick_to_bricks_pcFunction · 0.90

Tested by

no test coverage detected