import numpy as np from einops import rearrange from torchvision import transforms as tr statistics = { "vangogh": { "mean": [0.51376258, 0.48525076, 0.33816723], "std": [0.22747658, 0.2099791, 0.19311549] }, "monet": { "mean": [0.51991925, 0.51146052, 0.47215716], "std": [0.18806553, 0.17698975, 0.18903786] }, "cezanne": { "mean": [0.46261039, 0.4447964, 0.35517752], "std": [0.20604287, 0.18698825, 0.18653054] }, "photo": { "mean": [0.41229027, 0.40928956, 0.39267376], "std": [0.22338789, 0.2017264, 0.21975849] } } def get_transforms(name): mean, std = statistics[name]["mean"], statistics[name]["std"] val_transform = tr.Compose([ # tr.ToPILImage(), tr.Resize(size=(512, 512)), tr.ToTensor(), tr.Normalize(mean=mean, std=std), ]) def de_normalize(image, normalized=True): image = image.detach().cpu().numpy() if not normalized: return image image = rearrange(image, "c h w -> h w c") image = image * std + mean return np.clip(image, 0, 1) return val_transform, de_normalize def tensor_to_image(tensor, de_norm=None): tensor = tensor.squeeze(0) if de_norm is not None: tensor = de_norm(tensor) tensor = (tensor * 255).astype(np.uint8) return tensor