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