resubmit / data /Dataset.py
Orannue's picture
Upload 21 files
7501dfb verified
Raw History Blame Contribute Delete
3.59 kB
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