Team Ai
Apppublic

Dynamatrix/DiffBIR-OpenXLab

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
diffjpeg.py493 linesDownload Raw Back to image
1# https://github.com/XPixelGroup/BasicSR/blob/master/basicsr/utils/diffjpeg.py2"""3Modified from https://github.com/mlomnitz/DiffJPEG4 5For images not divisible by 86https://dsp.stackexchange.com/questions/35339/jpeg-dct-padding/35343#353437"""8import itertools9import numpy as np10import torch11import torch.nn as nn12from torch.nn import functional as F13 14# ------------------------ utils ------------------------#15y_table = np.array(16    [[16, 11, 10, 16, 24, 40, 51, 61], [12, 12, 14, 19, 26, 58, 60, 55], [14, 13, 16, 24, 40, 57, 69, 56],17     [14, 17, 22, 29, 51, 87, 80, 62], [18, 22, 37, 56, 68, 109, 103, 77], [24, 35, 55, 64, 81, 104, 113, 92],18     [49, 64, 78, 87, 103, 121, 120, 101], [72, 92, 95, 98, 112, 100, 103, 99]],19    dtype=np.float32).T20y_table = nn.Parameter(torch.from_numpy(y_table))21c_table = np.empty((8, 8), dtype=np.float32)22c_table.fill(99)23c_table[:4, :4] = np.array([[17, 18, 24, 47], [18, 21, 26, 66], [24, 26, 56, 99], [47, 66, 99, 99]]).T24c_table = nn.Parameter(torch.from_numpy(c_table))25 26 27def diff_round(x):28    """ Differentiable rounding function29    """30    return torch.round(x) + (x - torch.round(x))**331 32 33def quality_to_factor(quality):34    """ Calculate factor corresponding to quality35 36    Args:37        quality(float): Quality for jpeg compression.38 39    Returns:40        float: Compression factor.41    """42    if quality < 50:43        quality = 5000. / quality44    else:45        quality = 200. - quality * 246    return quality / 100.47 48 49# ------------------------ compression ------------------------#50class RGB2YCbCrJpeg(nn.Module):51    """ Converts RGB image to YCbCr52    """53 54    def __init__(self):55        super(RGB2YCbCrJpeg, self).__init__()56        matrix = np.array([[0.299, 0.587, 0.114], [-0.168736, -0.331264, 0.5], [0.5, -0.418688, -0.081312]],57                          dtype=np.float32).T58        self.shift = nn.Parameter(torch.tensor([0., 128., 128.]))59        self.matrix = nn.Parameter(torch.from_numpy(matrix))60 61    def forward(self, image):62        """63        Args:64            image(Tensor): batch x 3 x height x width65 66        Returns:67            Tensor: batch x height x width x 368        """69        image = image.permute(0, 2, 3, 1)70        result = torch.tensordot(image, self.matrix, dims=1) + self.shift71        return result.view(image.shape)72 73 74class ChromaSubsampling(nn.Module):75    """ Chroma subsampling on CbCr channels76    """77 78    def __init__(self):79        super(ChromaSubsampling, self).__init__()80 81    def forward(self, image):82        """83        Args:84            image(tensor): batch x height x width x 385 86        Returns:87            y(tensor): batch x height x width88            cb(tensor): batch x height/2 x width/289            cr(tensor): batch x height/2 x width/290        """91        image_2 = image.permute(0, 3, 1, 2).clone()92        cb = F.avg_pool2d(image_2[:, 1, :, :].unsqueeze(1), kernel_size=2, stride=(2, 2), count_include_pad=False)93        cr = F.avg_pool2d(image_2[:, 2, :, :].unsqueeze(1), kernel_size=2, stride=(2, 2), count_include_pad=False)94        cb = cb.permute(0, 2, 3, 1)95        cr = cr.permute(0, 2, 3, 1)96        return image[:, :, :, 0], cb.squeeze(3), cr.squeeze(3)97 98 99class BlockSplitting(nn.Module):100    """ Splitting image into patches101    """102 103    def __init__(self):104        super(BlockSplitting, self).__init__()105        self.k = 8106 107    def forward(self, image):108        """109        Args:110            image(tensor): batch x height x width111 112        Returns:113            Tensor:  batch x h*w/64 x h x w114        """115        height, _ = image.shape[1:3]116        batch_size = image.shape[0]117        image_reshaped = image.view(batch_size, height // self.k, self.k, -1, self.k)118        image_transposed = image_reshaped.permute(0, 1, 3, 2, 4)119        return image_transposed.contiguous().view(batch_size, -1, self.k, self.k)120 121 122class DCT8x8(nn.Module):123    """ Discrete Cosine Transformation124    """125 126    def __init__(self):127        super(DCT8x8, self).__init__()128        tensor = np.zeros((8, 8, 8, 8), dtype=np.float32)129        for x, y, u, v in itertools.product(range(8), repeat=4):130            tensor[x, y, u, v] = np.cos((2 * x + 1) * u * np.pi / 16) * np.cos((2 * y + 1) * v * np.pi / 16)131        alpha = np.array([1. / np.sqrt(2)] + [1] * 7)132        self.tensor = nn.Parameter(torch.from_numpy(tensor).float())133        self.scale = nn.Parameter(torch.from_numpy(np.outer(alpha, alpha) * 0.25).float())134 135    def forward(self, image):136        """137        Args:138            image(tensor): batch x height x width139 140        Returns:141            Tensor: batch x height x width142        """143        image = image - 128144        result = self.scale * torch.tensordot(image, self.tensor, dims=2)145        result.view(image.shape)146        return result147 148 149class YQuantize(nn.Module):150    """ JPEG Quantization for Y channel151 152    Args:153        rounding(function): rounding function to use154    """155 156    def __init__(self, rounding):157        super(YQuantize, self).__init__()158        self.rounding = rounding159        self.y_table = y_table160 161    def forward(self, image, factor=1):162        """163        Args:164            image(tensor): batch x height x width165 166        Returns:167            Tensor: batch x height x width168        """169        if isinstance(factor, (int, float)):170            image = image.float() / (self.y_table * factor)171        else:172            b = factor.size(0)173            table = self.y_table.expand(b, 1, 8, 8) * factor.view(b, 1, 1, 1)174            image = image.float() / table175        image = self.rounding(image)176        return image177 178 179class CQuantize(nn.Module):180    """ JPEG Quantization for CbCr channels181 182    Args:183        rounding(function): rounding function to use184    """185 186    def __init__(self, rounding):187        super(CQuantize, self).__init__()188        self.rounding = rounding189        self.c_table = c_table190 191    def forward(self, image, factor=1):192        """193        Args:194            image(tensor): batch x height x width195 196        Returns:197            Tensor: batch x height x width198        """199        if isinstance(factor, (int, float)):200            image = image.float() / (self.c_table * factor)201        else:202            b = factor.size(0)203            table = self.c_table.expand(b, 1, 8, 8) * factor.view(b, 1, 1, 1)204            image = image.float() / table205        image = self.rounding(image)206        return image207 208 209class CompressJpeg(nn.Module):210    """Full JPEG compression algorithm211 212    Args:213        rounding(function): rounding function to use214    """215 216    def __init__(self, rounding=torch.round):217        super(CompressJpeg, self).__init__()218        self.l1 = nn.Sequential(RGB2YCbCrJpeg(), ChromaSubsampling())219        self.l2 = nn.Sequential(BlockSplitting(), DCT8x8())220        self.c_quantize = CQuantize(rounding=rounding)221        self.y_quantize = YQuantize(rounding=rounding)222 223    def forward(self, image, factor=1):224        """225        Args:226            image(tensor): batch x 3 x height x width227 228        Returns:229            dict(tensor): Compressed tensor with batch x h*w/64 x 8 x 8.230        """231        y, cb, cr = self.l1(image * 255)232        components = {'y': y, 'cb': cb, 'cr': cr}233        for k in components.keys():234            comp = self.l2(components[k])235            if k in ('cb', 'cr'):236                comp = self.c_quantize(comp, factor=factor)237            else:238                comp = self.y_quantize(comp, factor=factor)239 240            components[k] = comp241 242        return components['y'], components['cb'], components['cr']243 244 245# ------------------------ decompression ------------------------#246 247 248class YDequantize(nn.Module):249    """Dequantize Y channel250    """251 252    def __init__(self):253        super(YDequantize, self).__init__()254        self.y_table = y_table255 256    def forward(self, image, factor=1):257        """258        Args:259            image(tensor): batch x height x width260 261        Returns:262            Tensor: batch x height x width263        """264        if isinstance(factor, (int, float)):265            out = image * (self.y_table * factor)266        else:267            b = factor.size(0)268            table = self.y_table.expand(b, 1, 8, 8) * factor.view(b, 1, 1, 1)269            out = image * table270        return out271 272 273class CDequantize(nn.Module):274    """Dequantize CbCr channel275    """276 277    def __init__(self):278        super(CDequantize, self).__init__()279        self.c_table = c_table280 281    def forward(self, image, factor=1):282        """283        Args:284            image(tensor): batch x height x width285 286        Returns:287            Tensor: batch x height x width288        """289        if isinstance(factor, (int, float)):290            out = image * (self.c_table * factor)291        else:292            b = factor.size(0)293            table = self.c_table.expand(b, 1, 8, 8) * factor.view(b, 1, 1, 1)294            out = image * table295        return out296 297 298class iDCT8x8(nn.Module):299    """Inverse discrete Cosine Transformation300    """301 302    def __init__(self):303        super(iDCT8x8, self).__init__()304        alpha = np.array([1. / np.sqrt(2)] + [1] * 7)305        self.alpha = nn.Parameter(torch.from_numpy(np.outer(alpha, alpha)).float())306        tensor = np.zeros((8, 8, 8, 8), dtype=np.float32)307        for x, y, u, v in itertools.product(range(8), repeat=4):308            tensor[x, y, u, v] = np.cos((2 * u + 1) * x * np.pi / 16) * np.cos((2 * v + 1) * y * np.pi / 16)309        self.tensor = nn.Parameter(torch.from_numpy(tensor).float())310 311    def forward(self, image):312        """313        Args:314            image(tensor): batch x height x width315 316        Returns:317            Tensor: batch x height x width318        """319        image = image * self.alpha320        result = 0.25 * torch.tensordot(image, self.tensor, dims=2) + 128321        result.view(image.shape)322        return result323 324 325class BlockMerging(nn.Module):326    """Merge patches into image327    """328 329    def __init__(self):330        super(BlockMerging, self).__init__()331 332    def forward(self, patches, height, width):333        """334        Args:335            patches(tensor) batch x height*width/64, height x width336            height(int)337            width(int)338 339        Returns:340            Tensor: batch x height x width341        """342        k = 8343        batch_size = patches.shape[0]344        image_reshaped = patches.view(batch_size, height // k, width // k, k, k)345        image_transposed = image_reshaped.permute(0, 1, 3, 2, 4)346        return image_transposed.contiguous().view(batch_size, height, width)347 348 349class ChromaUpsampling(nn.Module):350    """Upsample chroma layers351    """352 353    def __init__(self):354        super(ChromaUpsampling, self).__init__()355 356    def forward(self, y, cb, cr):357        """358        Args:359            y(tensor): y channel image360            cb(tensor): cb channel361            cr(tensor): cr channel362 363        Returns:364            Tensor: batch x height x width x 3365        """366 367        def repeat(x, k=2):368            height, width = x.shape[1:3]369            x = x.unsqueeze(-1)370            x = x.repeat(1, 1, k, k)371            x = x.view(-1, height * k, width * k)372            return x373 374        cb = repeat(cb)375        cr = repeat(cr)376        return torch.cat([y.unsqueeze(3), cb.unsqueeze(3), cr.unsqueeze(3)], dim=3)377 378 379class YCbCr2RGBJpeg(nn.Module):380    """Converts YCbCr image to RGB JPEG381    """382 383    def __init__(self):384        super(YCbCr2RGBJpeg, self).__init__()385 386        matrix = np.array([[1., 0., 1.402], [1, -0.344136, -0.714136], [1, 1.772, 0]], dtype=np.float32).T387        self.shift = nn.Parameter(torch.tensor([0, -128., -128.]))388        self.matrix = nn.Parameter(torch.from_numpy(matrix))389 390    def forward(self, image):391        """392        Args:393            image(tensor): batch x height x width x 3394 395        Returns:396            Tensor: batch x 3 x height x width397        """398        result = torch.tensordot(image + self.shift, self.matrix, dims=1)399        return result.view(image.shape).permute(0, 3, 1, 2)400 401 402class DeCompressJpeg(nn.Module):403    """Full JPEG decompression algorithm404 405    Args:406        rounding(function): rounding function to use407    """408 409    def __init__(self, rounding=torch.round):410        super(DeCompressJpeg, self).__init__()411        self.c_dequantize = CDequantize()412        self.y_dequantize = YDequantize()413        self.idct = iDCT8x8()414        self.merging = BlockMerging()415        self.chroma = ChromaUpsampling()416        self.colors = YCbCr2RGBJpeg()417 418    def forward(self, y, cb, cr, imgh, imgw, factor=1):419        """420        Args:421            compressed(dict(tensor)): batch x h*w/64 x 8 x 8422            imgh(int)423            imgw(int)424            factor(float)425 426        Returns:427            Tensor: batch x 3 x height x width428        """429        components = {'y': y, 'cb': cb, 'cr': cr}430        for k in components.keys():431            if k in ('cb', 'cr'):432                comp = self.c_dequantize(components[k], factor=factor)433                height, width = int(imgh / 2), int(imgw / 2)434            else:435                comp = self.y_dequantize(components[k], factor=factor)436                height, width = imgh, imgw437            comp = self.idct(comp)438            components[k] = self.merging(comp, height, width)439            #440        image = self.chroma(components['y'], components['cb'], components['cr'])441        image = self.colors(image)442 443        image = torch.min(255 * torch.ones_like(image), torch.max(torch.zeros_like(image), image))444        return image / 255445 446 447# ------------------------ main DiffJPEG ------------------------ #448 449 450class DiffJPEG(nn.Module):451    """This JPEG algorithm result is slightly different from cv2.452    DiffJPEG supports batch processing.453 454    Args:455        differentiable(bool): If True, uses custom differentiable rounding function, if False, uses standard torch.round456    """457 458    def __init__(self, differentiable=True):459        super(DiffJPEG, self).__init__()460        if differentiable:461            rounding = diff_round462        else:463            rounding = torch.round464 465        self.compress = CompressJpeg(rounding=rounding)466        self.decompress = DeCompressJpeg(rounding=rounding)467 468    def forward(self, x, quality):469        """470        Args:471            x (Tensor): Input image, bchw, rgb, [0, 1]472            quality(float): Quality factor for jpeg compression scheme.473        """474        factor = quality475        if isinstance(factor, (int, float)):476            factor = quality_to_factor(factor)477        else:478            for i in range(factor.size(0)):479                factor[i] = quality_to_factor(factor[i])480        h, w = x.size()[-2:]481        h_pad, w_pad = 0, 0482        # why should use 16483        if h % 16 != 0:484            h_pad = 16 - h % 16485        if w % 16 != 0:486            w_pad = 16 - w % 16487        x = F.pad(x, (0, w_pad, 0, h_pad), mode='constant', value=0)488 489        y, cb, cr = self.compress(x, factor=factor)490        recovered = self.decompress(y, cb, cr, (h + h_pad), (w + w_pad), factor=factor)491        recovered = recovered[:, :, 0:h, 0:w]492        return recovered493