Spaces:
Running
Running
| 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 | |