MCPcopy Create free account
hub / github.com/MaureenZOU/TSAM / _evaluate_data_loader

Method _evaluate_data_loader

src/trainer/inference.py:141–254  ·  view source on GitHub ↗
(self, epoch=None, output_root_dir=None, data_loader=None, name='test')

Source from the content-addressed store, hash-verified

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

Callers 1

evaluate_test_setMethod · 0.95

Calls 12

get_evaluate_indexMethod · 0.95
index_dataMethod · 0.95
_process_dataMethod · 0.95
_unpack_dataMethod · 0.95
_store_gated_valuesMethod · 0.95
_evaluate_test_videoMethod · 0.95
_write_imagesMethod · 0.95
get_i3d_activationsFunction · 0.90
save_frames_to_dirFunction · 0.90
get_fid_scoreFunction · 0.90
appendMethod · 0.80
set_stepMethod · 0.80

Tested by

no test coverage detected