CoolFace
Apppublic

DimaTivator/cycle-gan

sourceHugging Facemitupdated 1y agoView on Hugging Face
1likes
preprocessing.py55 linesDownload Raw Back to root
1import numpy as np2from einops import rearrange3from torchvision import transforms as tr4 5 6statistics = {7    "vangogh": {8        "mean": [0.51376258, 0.48525076, 0.33816723],9        "std": [0.22747658, 0.2099791, 0.19311549]10    },11    "monet": {12        "mean": [0.51991925, 0.51146052, 0.47215716],13        "std": [0.18806553, 0.17698975, 0.18903786]14    },15    "cezanne": {16        "mean": [0.46261039, 0.4447964, 0.35517752],17        "std": [0.20604287, 0.18698825, 0.18653054]18    },19    "photo": {20        "mean": [0.41229027, 0.40928956, 0.39267376],21        "std": [0.22338789, 0.2017264, 0.21975849]22    }23}24 25 26def get_transforms(name):27    mean, std = statistics[name]["mean"], statistics[name]["std"]28    29    val_transform = tr.Compose([30        # tr.ToPILImage(),31        tr.Resize(size=(512, 512)),32        tr.ToTensor(),33        tr.Normalize(mean=mean, std=std),34    ])35    36    def de_normalize(image, normalized=True):37        image = image.detach().cpu().numpy()38        39        if not normalized:40            return image41        42        image = rearrange(image, "c h w -> h w c")43        image = image * std + mean44        return np.clip(image, 0, 1)45    46    return val_transform, de_normalize47 48 49def tensor_to_image(tensor, de_norm=None):50    tensor = tensor.squeeze(0)51    if de_norm is not None:52        tensor = de_norm(tensor)53    tensor = (tensor * 255).astype(np.uint8) 54    return tensor55