"""Definition of Visual Critial Discourse Analysis dataloader.""" import random from PIL import Image from torch import Tensor from torch.utils.data import Dataset from torchvision.transforms import ColorJitter, InterpolationMode, Normalize from torchvision.transforms.functional import ( hflip, pad, resize, rotate, to_pil_image, to_tensor, ) from shared.repositories import ImageRepository # resnet18 original normalization values RESNET_NORMALIZE_MEAN = [0.485, 0.456, 0.406] RESNET_NORMALIZE_STD = [0.229, 0.224, 0.225] class VCDADataset(Dataset): """VCDA dataset class.""" def __init__( self, data_name_list: list[str], do_augment: bool = False, random_annotations: bool = False, ): super().__init__() self.data_name_list = data_name_list self.do_augment = do_augment self.random_annotations = random_annotations self.normalize_mean = RESNET_NORMALIZE_MEAN self.normalize_std = RESNET_NORMALIZE_STD # prepare augmentation functions self.normalize = Normalize( mean=self.normalize_mean, std=self.normalize_std, ) self.color_jitter = ColorJitter( brightness=1e-1, contrast=8e-2, saturation=8e-2, ) def __len__(self): return len(self.data_name_list) def __getitem__(self, idx): # get image from database image_name = self.data_name_list[idx] with ImageRepository() as repo: image_data = repo.get_data(image_name) assert image_data is not None tensor = self.image_to_tensor(image_data.image) if self.do_augment: tensor = self.augment(tensor) return tensor def image_to_tensor( self, image: Image.Image, ) -> Tensor: """Load images tensor from bytes.""" tensor = to_tensor(image) tensor /= 255 # normalize 8-bit image tensor = self.square_pad(tensor=tensor) tensor = resize( img=tensor, size=(512, 512), interpolation=InterpolationMode.BICUBIC, ) return tensor @staticmethod def square_pad( tensor: Tensor, ) -> Tensor: """Pads image to a square with side length equal to the largest side of the input image.""" assert isinstance(tensor, Tensor) # B, nc, w, h = img.shape h = tensor.shape[-2] w = tensor.shape[-1] if h == w: return tensor max_wh = max([h, w]) hp = int((max_wh - w) / 2) vp = int((max_wh - h) / 2) padding = (hp, vp, hp, vp) tensor = pad(tensor, padding, 0, 'constant') return tensor def augment( self, tensor: Tensor, ) -> Tensor: """Augment image with random horizontal flips, rotations and color jitter.""" assert isinstance(tensor, Tensor) # left-right flip if random.random() >= 0.5: tensor = hflip(tensor) # rotation rnd = random.random() if rnd < 0.25: tensor = rotate(tensor, angle=90) if rnd < 0.5: tensor = rotate(tensor, angle=180) if rnd < 0.75: tensor = rotate(tensor, angle=270) # color jitter tensor = self.color_jitter(tensor) return tensor def reverse_normalise( self, tensor: Tensor, ) -> Image.Image: """Reverse normalization to get an image that can be interpreted by humans.""" assert isinstance(tensor, Tensor) tensor *= Tensor(self.normalize_std).reshape((3, 1, 1)) tensor += Tensor(self.normalize_mean).reshape((3, 1, 1)) image = to_pil_image( pic=tensor, mode='RGB', ) return image