(
result_images, gt_images, masks,
flownet_checkpoint_path, evaluate_warping_error=False,
printlog=True
)
| 166 | |
| 167 | |
| 168 | def 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)}") |
no test coverage detected