File size: 3,231 Bytes
ed46d32
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
import torch
from torch import nn
from torch.utils.data import DataLoader
from torchvision import transforms
from torchvision import transforms as T
from torchvision.datasets import Imagenette

MEAN = [0.485, 0.456, 0.406]
STD = [0.229, 0.224, 0.225]

# https://gist.github.com/yrevar/942d3a0ac09ec9e5eb3a "text: imagenet 1000 class idx to human readable labels (Fox, E ..."
IMAGENETTE_TO_IMAGENET = {
    0: 0,  # tench
    1: 217,  # English springer
    2: 482,  # cassette player
    3: 491,  # chain saw
    4: 497,  # church
    5: 566,  # French horn
    6: 569,  # garbage truck
    7: 571,  # gas pump
    8: 574,  # golf ball
    9: 701,  # parachute
}


def imagenette_label_to_imagenet(label):
    return IMAGENETTE_TO_IMAGENET[label]


def my_normalize():
    return T.Normalize(
        mean=[0.5, 0.5, 0.5],
        std=[0.5, 0.5, 0.5],
    )


def my_denormalize(x):
    return (x + 1) / 2  # Map from [-1,1] to [0,1]


class FromMyNormalizeToImageNet(nn.Module):
    def __init__(self):
        super().__init__()
        self.imagenet_norm = T.Normalize(
            mean=MEAN,
            std=STD,
        )

    def forward(self, x):
        x = (x + 1) / 2  # Map from [-1,1] to [0,1]
        return self.imagenet_norm(x)


def get_transform():
    return transforms.Compose(
        [
            transforms.Resize(256),
            transforms.CenterCrop(224),
            transforms.ToTensor(),
            my_normalize(),
        ]
    )


class ConditionalTransform(nn.Module):
    def __init__(self):
        super().__init__()
        self.resize_224 = T.Resize(224)
        self.resize_256 = T.Resize(256)
        self.crop_224 = T.CenterCrop(224)
        self.to_tensor = T.ToTensor()
        self.normalize = my_normalize()

    def __call__(self, img):
        # img: PIL.Image
        w, h = img.size
        if w == h:
            img = self.resize_224(img)
        else:
            img = self.resize_256(img)
            img = self.crop_224(img)
        img = self.to_tensor(img)
        img = self.normalize(img)
        return img


def get_dataset(download=False):
    return Imagenette(
        root="./data",
        split="val",  # or "train"
        size="160px",  # can also be "320" or "full"
        download=download,
        transform=get_transform(),
        target_transform=imagenette_label_to_imagenet,
    )


def get_examples(loader, all_classes=False):
    class_indices = list(range(10)) if all_classes else [0, 2, 4, 6, 8]
    target_classes = [imagenette_label_to_imagenet(l) for l in class_indices]

    selected = {}
    seen_classes = set()

    for batch in loader:
        images, labels = batch
        for img, label in zip(images, labels):
            label = label.item()
            if label in target_classes and label not in seen_classes:
                selected[label] = img.unsqueeze(0)
                seen_classes.add(label)
            if len(seen_classes) == len(target_classes):
                break
        if len(seen_classes) == len(target_classes):
            break

    # Concat images
    images = torch.cat([selected[c] for c in target_classes], dim=0)  # [5, C, H, W]
    labels = torch.tensor(target_classes)  # [5]

    return images, labels