(opt, model, set_path, meters, is_subm=False, subm_cbricks=None, subm_meters=None)
| 120 | |
| 121 | |
| 122 | def 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: |
no test coverage detected