(self, epoch=None, output_root_dir=None, data_loader=None, name='test')
| 139 | return input_data |
| 140 | |
| 141 | def _evaluate_data_loader(self, epoch=None, output_root_dir=None, data_loader=None, name='test'): |
| 142 | total_length = 0 |
| 143 | total_warp_error = 0 if self.evaluate_test_warp_error else None |
| 144 | total_error = 0 |
| 145 | total_psnr = 0 |
| 146 | total_ssim = 0 |
| 147 | total_p_dist = 0 |
| 148 | |
| 149 | if output_root_dir is None: |
| 150 | output_root_dir = self.test_output_root_dir |
| 151 | val_log_dir = os.path.join(self.checkpoint_dir, 'val.log') |
| 152 | if epoch is not None: |
| 153 | output_root_dir = os.path.join(output_root_dir, f"epoch_{epoch}") |
| 154 | output_root_dir = os.path.join(output_root_dir, name) |
| 155 | |
| 156 | output_i3d_activations = [] |
| 157 | real_i3d_activations = [] |
| 158 | with torch.no_grad(): |
| 159 | for batch_idx, data in enumerate(data_loader): |
| 160 | n,t,c,h,w = data['input_tensors'].shape |
| 161 | start, end, iter_start, iter_end, out_start, out_end = self.get_evaluate_index(t, vid_len=self.valid_length, valid_inter=self.valid_interval) |
| 162 | |
| 163 | inputs_ = torch.zeros((t,c,h,w)) |
| 164 | outputs_ = torch.zeros((t,c,h,w)) |
| 165 | targets_ = torch.zeros((t,c,h,w)) |
| 166 | masks_ = torch.zeros((t,1,h,w)) |
| 167 | |
| 168 | for m,n,i,j,p,q in zip(start, end, iter_start, iter_end, out_start, out_end): |
| 169 | input_ = self.index_data(data, i,j) |
| 170 | _, _, data_input, model_output = self._process_data(input_, validation=True) |
| 171 | inputs, outputs, targets, masks = self._unpack_data(data_input, model_output) |
| 172 | if self.store_gated_values: |
| 173 | out_dir = os.path.join(output_root_dir, 'gated_values', f'input_{batch_idx:04}') |
| 174 | self._store_gated_values(out_dir) |
| 175 | outputs = outputs.clamp(0, 1) |
| 176 | |
| 177 | if self.evaluate_score: |
| 178 | # get i3d activation |
| 179 | output_i3d_activations.append(get_i3d_activations(outputs).cpu().numpy()) |
| 180 | real_i3d_activations.append(get_i3d_activations(targets).cpu().numpy()) |
| 181 | |
| 182 | assert len(outputs) == 1 # Batch size = 1 for testing |
| 183 | inputs = inputs[0].cpu() |
| 184 | outputs = outputs[0].cpu() |
| 185 | targets = targets[0].cpu() |
| 186 | masks = masks[0].cpu() |
| 187 | |
| 188 | inputs_[p:q,:,:,:] = inputs[m:n,:,:,:] |
| 189 | outputs_[p:q,:,:,:] = outputs[m:n,:,:,:] |
| 190 | targets_[p:q,:,:,:] = targets[m:n,:,:,:] |
| 191 | masks_[p:q,:,:,:] = masks[m:n,:,:,:] |
| 192 | |
| 193 | # if epoch is not None and epoch == 0: |
| 194 | # # Save inputs to output_dir |
| 195 | # output_dir = os.path.join(output_root_dir, 'inputs', f"input_{batch_idx:04}") |
| 196 | # self.logger.debug(f"Saving batch {batch_idx} input to {output_dir}") |
| 197 | # save_frames_to_dir([self.toPILImage(t) for t in inputs.cpu()], output_dir) |
| 198 |
no test coverage detected