Spaces:
Sleeping
Sleeping
File size: 1,871 Bytes
77b48d9 | 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 | from torch.utils.data import Dataset
import os
from PIL import Image
from torchvision import transforms
class ImageFolderDataset(Dataset):
def __init__(self, root, transform=None):
super(ImageFolderDataset,self).__init__()
self.root=root
self.transform=transform
self.files=list(os.listdir(root))
self.files=[p for p in self.files if p.endswith(('.jpg','.png','.jpeg'))]
def __len__(self):
return len(self.files)
def __getitem__(self,idx):
image_path=os.path.join(self.root,self.files[idx])
image=Image.open(image_path).convert('RGB')
if self.transform:
image=self.transform(image)
return image
def get_transform(size,crop,final_size):
transform_list=[]
if size>0:
transform_list.append(transforms.Resize(size))
if crop:
transform_list.append(transforms.RandomCrop(final_size))
else:
transform_list.append(transforms.Resize(final_size))
transform_list.append(transforms.ToTensor())
return transforms.Compose(transform_list)
def adaptive_instance_normalization(content_feat, style_feat):
size=content_feat.size()
style_mean, style_std=calc_mean_std(style_feat)
content_mean, content_std= calc_mean_std(content_feat)
normalized_content_feat = (content_feat - content_mean.expand(size)) / content_std.expand(size)
return normalized_content_feat*style_std.expand(size)+style_mean.expand(size)
def calc_mean_std(feat, eps=1e-5):
size=feat.size()
assert(len(size)==4)
batch_size, channels=size[:2]
feat_mean=feat.view(batch_size,channels,-1).mean(dim=2).view(batch_size,channels,1,1)
feat_var = feat.view(batch_size, channels, -1).var(dim=2, unbiased=False) + eps
feat_std = feat_var.sqrt().view(batch_size, channels, 1, 1)
return feat_mean, feat_std |