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

Function evaluate_video_error

src/evaluate.py:168–231  ·  view source on GitHub ↗
(
    result_images, gt_images, masks,
    flownet_checkpoint_path, evaluate_warping_error=False,
    printlog=True
)

Source from the content-addressed store, hash-verified

166
167
168def evaluate_video_error(
169 result_images, gt_images, masks,
170 flownet_checkpoint_path, evaluate_warping_error=False,
171 printlog=True
172):
173 total_error = 0
174 total_psnr = 0
175 total_ssim = 0
176 total_p_dist = 0
177 # worker = torchvision.transforms.Resize(SIZE, Image.BILINEAR)
178
179 for i, (result, gt, mask) in enumerate(
180 zip(result_images, gt_images, masks)
181 ):
182 # mask = np.expand_dims(mask, 2)
183 # gt = worker(gt)
184 # mask = worker(mask)
185
186 mse, ssim_value, psnr_value, p_dist = evaluate_image(gt, result)
187 total_error += mse
188 total_ssim += ssim_value
189 total_psnr += psnr_value
190 total_p_dist += p_dist
191 logger.debug(
192 f"Frame {i}: MSE {mse} PSNR {psnr_value} SSIM {ssim_value} "
193 f"Percep. Dist. {p_dist}"
194 )
195
196 if evaluate_warping_error:
197 init_warping_model(flownet_checkpoint_path)
198 # These readers are lists of images
199 # After np.array, they are in shape (H, W, C)
200 # While the required input of the temporal_warping_error is (B, L, C, H, W)
201 # So the tensors are unsqueezed and permuted into such shape
202 targets = torch.Tensor(
203 [np.array(x) for x in gt_images]
204 ).unsqueeze(0).permute(0, 1, 4, 2, 3)
205 masks = torch.Tensor(
206 [np.array(x) for x in masks]
207 ).unsqueeze(3).unsqueeze(0).permute(0, 1, 4, 2, 3)
208 outputs = torch.Tensor([np.array(x) for x in result_images]).unsqueeze(0).permute(0, 1, 4, 2, 3)
209 data_input = {
210 "targets": targets,
211 "masks": masks
212 }
213 model_output = {"outputs": outputs}
214 warping_error = temporal_warping_error(data_input, model_output).cpu().item()
215 if printlog:
216 logger.info(f"Warping error: {warping_error}")
217 else:
218 warping_error = 0
219
220 if printlog:
221 logger.info(f"Avg MSE: {total_error / len(result_images)}")
222 logger.info(f"Avg PSNR: {total_psnr / len(result_images)}")
223 logger.info(f"Avg SSIM: {total_ssim / len(result_images)}")
224 logger.info(
225 f"Avg Perce. Dist.: {total_p_dist / len(result_images)}")

Callers 3

_evaluate_test_videoMethod · 0.90
_evaluate_test_videoMethod · 0.90
evaluate_videoFunction · 0.85

Calls 2

evaluate_imageFunction · 0.85
init_warping_modelFunction · 0.85

Tested by

no test coverage detected