DimaTivator/cycle-gan
1
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 