File size: 3,594 Bytes
7501dfb | 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 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 | 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
|