Spaces:
Running
Running
File size: 1,420 Bytes
c3b5390 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 | 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
|