import os.path from PIL import Image from torch.utils.data import Dataset from torchvision import transforms import os import os.path import torch from data.crop import * from data.transforms import Global_crops, dino_structure_transforms, dino_texture_transforms class SingleImageDataset(Dataset): def __init__(self, cfg,path1,path2,mode='global',mask_path='1'): self.cfg = cfg self.structure_transforms = dino_structure_transforms if cfg['use_augmentations'] else transforms.Compose([]) self.texture_transforms = dino_texture_transforms if cfg['use_augmentations'] else transforms.Compose([]) self.base_transform = transforms.Compose([ transforms.ToTensor(), ]) self.global_A_patches = transforms.Compose( [ self.structure_transforms, Global_crops(n_crops=cfg['global_A_crops_n_crops'], min_cover=cfg['global_A_crops_min_cover'], last_transform=self.base_transform) ] ) self.global_B_patches = transforms.Compose( [ self.texture_transforms, Global_crops(n_crops=cfg['global_B_crops_n_crops'], min_cover=cfg['global_B_crops_min_cover'], last_transform=self.base_transform) ] ) # open images # dir_A = os.path.join(cfg['dataroot'], 'A') # dir_B = os.path.join(cfg['dataroot'], 'B') A_path = path1 B_path = path2 if mode=='black': self.A_img = crop1(A_path,mask_path) self.B_img = crop1(B_path,mask_path) elif mode=='local': self.A_img = crop3(A_path,mask_path) self.B_img = crop3(B_path,mask_path) else: self.A_img = Image.open(A_path).convert('RGB') self.B_img = Image.open(B_path).convert('RGB') if cfg['A_resize'] > 0: self.A_img = transforms.Resize(cfg['A_resize'])(self.A_img) if cfg['B_resize'] > 0: self.B_img = transforms.Resize(cfg['B_resize'])(self.B_img) if cfg['direction'] == 'BtoA': self.A_img, self.B_img = self.B_img, self.A_img print("Image sizes %s and %s" % (str(self.A_img.size), str(self.B_img.size))) self.step = torch.zeros(1) - 1 def get_A(self): return self.base_transform(self.A_img).unsqueeze(0) def __getitem__(self, index): self.step += 1 sample = {'step': self.step} if self.step % self.cfg['entire_A_every'] == 0: sample['A'] = self.get_A() sample['A_global'] = self.global_A_patches(self.A_img) sample['B_global'] = self.global_B_patches(self.B_img) return sample def calculate_global_ssim_loss(self): loss = 0.0 sample['A_global'] = self.global_A_patches(self.A_img) sample['B_global'] = self.global_B_patches(self.B_img) inputs=sample['A_global'] outputs=sample['B_global'] for a, b in zip(inputs, outputs): # avoid memory limitations a = self.global_transform(a) b = self.global_transform(b) with torch.no_grad(): target_keys_self_sim = self.extractor.get_keys_self_sim_from_input(a.unsqueeze(0), layer_num=11) keys_ssim = self.extractor.get_keys_self_sim_from_input(b.unsqueeze(0), layer_num=11) loss += F.mse_loss(keys_ssim, target_keys_self_sim) return loss def __len__(self): return 1