(image)
| 26 | return 20 * torch.log10(1.0 / torch.sqrt(mse)) |
| 27 | |
| 28 | def gradient_map(image): |
| 29 | sobel_x = torch.tensor([[-1, 0, 1], [-2, 0, 2], [-1, 0, 1]]).float().unsqueeze(0).unsqueeze(0).cuda()/4 |
| 30 | sobel_y = torch.tensor([[-1, -2, -1], [0, 0, 0], [1, 2, 1]]).float().unsqueeze(0).unsqueeze(0).cuda()/4 |
| 31 | |
| 32 | grad_x = torch.cat([F.conv2d(image[i].unsqueeze(0), sobel_x, padding=1) for i in range(image.shape[0])]) |
| 33 | grad_y = torch.cat([F.conv2d(image[i].unsqueeze(0), sobel_y, padding=1) for i in range(image.shape[0])]) |
| 34 | magnitude = torch.sqrt(grad_x ** 2 + grad_y ** 2) |
| 35 | magnitude = magnitude.norm(dim=0, keepdim=True) |
| 36 | |
| 37 | return magnitude |
| 38 | |
| 39 | def colormap(map, cmap="turbo"): |
| 40 | colors = torch.tensor(plt.cm.get_cmap(cmap).colors).to(map.device) |
no test coverage detected