File size: 2,136 Bytes
9e14838
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from typing import Callable

import numpy as np
from torchvision.datasets import MNIST

from ..config import Config
from ..dataset import BaseDataModule


class MNISTDataset(MNIST):
    def __init__(
        self,
        train: bool = True,
        preprocess: None | Callable = None,
        augmentations: None | Callable = None,
    ):
        super().__init__(root="datasets/other/", train=train, download=True)
        self.preprocess = preprocess
        self.augmentations = augmentations

    def __getitem__(self, idx):
        image, label = super().__getitem__(idx)
        if self.augmentations is not None:
            image = self.augmentations(image)
        if self.preprocess is not None:
            image = self.preprocess(image)
        return {
            "image": image,
            "label": label,
            "path": f"{idx}: {label}",
            "idx": idx,
        }

    def print_statistics(self):
        print(f"Number of samples: {len(self)}")
        unique, counts = np.unique(self.targets, return_counts=True)
        print("Class distribution")
        names = self.get_class_names()
        for u, c in zip(unique, counts):
            print(f"Class {u} ({names[u]}): {c}")

    def get_class_names(self) -> dict[int, str]:
        return {i: str(i) for i in range(10)}


class MNISTDataModule(BaseDataModule):
    def __init__(self, config: Config, preprocess: None | Callable = None):
        super().__init__(config, preprocess)

    def setup(self, stage: str):
        # Initialize datasets
        if stage == "fit" or stage == "validate":
            self.train_dataset = MNISTDataset(train=True, preprocess=self.preprocess)
            self.val_dataset = MNISTDataset(train=False, preprocess=self.preprocess)

            print("\nTrain dataset")
            self.train_dataset.print_statistics()

            print("\nValidation dataset")
            self.val_dataset.print_statistics()

        if stage == "test":
            self.test_dataset = MNISTDataset(train=False, preprocess=self.preprocess)

            print("\nTest dataset")
            self.test_dataset.print_statistics()