GoodWin/Deep-Multi-scale
0
1import torch.utils.data as data2from PIL import Image3import torchvision.transforms as transforms4 5 6class BaseDataset(data.Dataset):7 def __init__(self):8 super(BaseDataset, self).__init__()9 10 def name(self):11 return 'BaseDataset'12 13 @staticmethod14 def modify_commandline_options(parser, is_train):15 return parser16 17 def initialize(self, opt):18 pass19 20 def __len__(self):21 return 022 23 24def get_transform(opt):25 transform_list = []26 if opt.resize_or_crop == 'resize_and_crop':27 28 osize = [opt.loadSize, opt.loadSize]29 transform_list.append(transforms.Resize(osize, Image.BICUBIC))30 transform_list.append(transforms.RandomCrop(opt.fineSize))31 elif opt.resize_or_crop == 'crop':32 transform_list.append(transforms.RandomCrop(opt.fineSize))33 elif opt.resize_or_crop == 'scale_width':34 transform_list.append(transforms.Lambda(35 lambda img: __scale_width(img, opt.fineSize)))36 elif opt.resize_or_crop == 'scale_width_and_crop':37 transform_list.append(transforms.Lambda(38 lambda img: __scale_width(img, opt.loadSize)))39 transform_list.append(transforms.RandomCrop(opt.fineSize))40 elif opt.resize_or_crop == 'none':41 transform_list.append(transforms.Lambda(42 lambda img: __adjust(img)))43 else:44 raise ValueError('--resize_or_crop %s is not a valid option.' % opt.resize_or_crop)45 46 if opt.isTrain and not opt.no_flip:47 transform_list.append(transforms.RandomHorizontalFlip())48 49 transform_list += [transforms.ToTensor(),50 transforms.Normalize((0.5, 0.5, 0.5),51 (0.5, 0.5, 0.5))]52 return transforms.Compose(transform_list)53 54# just modify the width and height to be multiple of 455def __adjust(img):56 ow, oh = img.size57 58 # the size needs to be a multiple of this number, 59 # because going through generator network may change img size60 # and eventually cause size mismatch error61 mult = 4 62 if ow % mult == 0 and oh % mult == 0:63 return img64 w = (ow - 1) // mult65 w = (w + 1) * mult66 h = (oh - 1) // mult67 h = (h + 1) * mult68 69 if ow != w or oh != h:70 __print_size_warning(ow, oh, w, h)71 72 return img.resize((w, h), Image.BICUBIC)73 74 75def __scale_width(img, target_width):76 ow, oh = img.size77 78 # the size needs to be a multiple of this number, 79 # because going through generator network may change img size80 # and eventually cause size mismatch error 81 mult = 482 assert target_width % mult == 0, "the target width needs to be multiple of %d." % mult83 if (ow == target_width and oh % mult == 0):84 return img85 w = target_width86 target_height = int(target_width * oh / ow)87 m = (target_height - 1) // mult88 h = (m + 1) * mult89 90 if target_height != h:91 __print_size_warning(target_width, target_height, w, h)92 93 return img.resize((w, h), Image.BICUBIC)94 95 96def __print_size_warning(ow, oh, w, h):97 if not hasattr(__print_size_warning, 'has_printed'):98 print("The image size needs to be a multiple of 4. "99 "The loaded image size was (%d, %d), so it was adjusted to "100 "(%d, %d). This adjustment will be done to all images "101 "whose sizes are not multiples of 4" % (ow, oh, w, h))102 __print_size_warning.has_printed = True103 104 105 