(img1, img2)
| 10 | |
| 11 | |
| 12 | def calculate_PSNR(img1, img2): |
| 13 | img1 = torch.clamp(img1, 0, 1) |
| 14 | img2 = torch.clamp(img2, 0, 1) |
| 15 | mse = torch.mean((img1 - img2) ** 2, dim=[1,2,3]) |
| 16 | psnrs = 20 * torch.log10(1 / torch.sqrt(mse)) |
| 17 | return psnrs.mean() |
| 18 | |
| 19 | |
| 20 | class ImageMetricTracker(nn.Module): |