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()
|