Dynamatrix/DiffBIR-OpenXLab
0
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 