File size: 2,310 Bytes
872b0a0 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 | import numpy as np
import torch
from PIL import ImageFilter
from torchvision import transforms as tf
def get_ap_transforms(cfg):
transforms = []
if cfg.cj:
transforms.append(
ColorJitter(
brightness=cfg.cj_bri,
contrast=cfg.cj_con,
saturation=cfg.cj_sat,
hue=cfg.cj_hue,
)
)
transforms.append(ToPILImage())
if cfg.gblur:
transforms.append(RandomGaussianBlur(0.5, 3))
transforms.append(ToTensor())
if cfg.gamma:
transforms.append(RandomGamma(min_gamma=0.7, max_gamma=1.5, clip_image=True))
return tf.Compose(transforms)
# from https://github.com/visinf/irr/blob/master/datasets/transforms.py
class ToPILImage(tf.ToPILImage):
def __call__(self, imgs):
return [super(ToPILImage, self).__call__(im) for im in imgs]
class ColorJitter(tf.ColorJitter):
def __call__(self, imgs):
_, h, w = imgs[0].shape
new_big_img = self.forward(torch.concatenate(imgs, dim=1))
return list(torch.split(new_big_img, h, dim=1))
# return [self.foward(im) for im in imgs]
class ToTensor(tf.ToTensor):
def __call__(self, imgs):
return [super(ToTensor, self).__call__(im) for im in imgs]
class RandomGamma:
def __init__(self, min_gamma=0.7, max_gamma=1.5, clip_image=False):
self._min_gamma = min_gamma
self._max_gamma = max_gamma
self._clip_image = clip_image
@staticmethod
def get_params(min_gamma, max_gamma):
return np.random.uniform(min_gamma, max_gamma)
@staticmethod
def adjust_gamma(image, gamma, clip_image):
adjusted = torch.pow(image, gamma)
if clip_image:
adjusted.clamp_(0.0, 1.0)
return adjusted
def __call__(self, imgs):
gamma = self.get_params(self._min_gamma, self._max_gamma)
return [self.adjust_gamma(im, gamma, self._clip_image) for im in imgs]
class RandomGaussianBlur:
def __init__(self, p, max_k_sz):
self.p = p
self.max_k_sz = max_k_sz
def __call__(self, imgs):
if np.random.random() < self.p:
radius = np.random.uniform(0, self.max_k_sz)
imgs = [im.filter(ImageFilter.GaussianBlur(radius)) for im in imgs]
return imgs
|