Team Ai
Apppublic

Dynamatrix/DiffBIR-OpenXLab

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
metrics.py67 linesDownload Raw Back to utils
1import torch2import lpips3 4from .image import rgb2ycbcr_pt5from .common import frozen_module6 7 8# https://github.com/XPixelGroup/BasicSR/blob/033cd6896d898fdd3dcda32e3102a792efa1b8f4/basicsr/metrics/psnr_ssim.py#L529def calculate_psnr_pt(img, img2, crop_border, test_y_channel=False):10    """Calculate PSNR (Peak Signal-to-Noise Ratio) (PyTorch version).11 12    Reference: https://en.wikipedia.org/wiki/Peak_signal-to-noise_ratio13 14    Args:15        img (Tensor): Images with range [0, 1], shape (n, 3/1, h, w).16        img2 (Tensor): Images with range [0, 1], shape (n, 3/1, h, w).17        crop_border (int): Cropped pixels in each edge of an image. These pixels are not involved in the calculation.18        test_y_channel (bool): Test on Y channel of YCbCr. Default: False.19 20    Returns:21        float: PSNR result.22    """23 24    assert img.shape == img2.shape, (f'Image shapes are different: {img.shape}, {img2.shape}.')25 26    if crop_border != 0:27        img = img[:, :, crop_border:-crop_border, crop_border:-crop_border]28        img2 = img2[:, :, crop_border:-crop_border, crop_border:-crop_border]29 30    if test_y_channel:31        img = rgb2ycbcr_pt(img, y_only=True)32        img2 = rgb2ycbcr_pt(img2, y_only=True)33 34    img = img.to(torch.float64)35    img2 = img2.to(torch.float64)36 37    mse = torch.mean((img - img2)**2, dim=[1, 2, 3])38    return 10. * torch.log10(1. / (mse + 1e-8))39 40 41class LPIPS:42    43    def __init__(self, net: str) -> None:44        self.model = lpips.LPIPS(net=net)45        frozen_module(self.model)46    47    @torch.no_grad()48    def __call__(self, img1: torch.Tensor, img2: torch.Tensor, normalize: bool) -> torch.Tensor:49        """50        Compute LPIPS.51        52        Args:53            img1 (torch.Tensor): The first image (NCHW, RGB, [-1, 1]). Specify `normalize` if input 54                image is range in [0, 1].55            img2 (torch.Tensor): The second image (NCHW, RGB, [-1, 1]). Specify `normalize` if input 56                image is range in [0, 1].57            normalize (bool): If specified, the input images will be normalized from [0, 1] to [-1, 1].58            59        Returns:60            lpips_values (torch.Tensor): The lpips scores of this batch.61        """62        return self.model(img1, img2, normalize=normalize)63    64    def to(self, device: str) -> "LPIPS":65        self.model.to(device)66        return self67