Team Ai
Apppublic

GoodWin/Deep-Multi-scale

sourceHugging Faceupdated 5y agoView on Hugging Face
0likes
base_dataset.py105 linesDownload Raw Back to data
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