File size: 1,437 Bytes
0c25286
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Data loading for chest X-ray classification.

Default layout is the Kaggle "Chest X-Ray Images (Pneumonia)" dataset
(Kermany et al.), an ImageFolder tree:

    data/
      train/{NORMAL,PNEUMONIA}/*.jpeg
      val/{NORMAL,PNEUMONIA}/*.jpeg
      test/{NORMAL,PNEUMONIA}/*.jpeg

Grayscale X-rays are expanded to 3 channels to match the ImageNet-pretrained
DenseNet stem.
"""
from torchvision import datasets, transforms
from torch.utils.data import DataLoader

IMAGENET_MEAN = [0.485, 0.456, 0.406]
IMAGENET_STD = [0.229, 0.224, 0.225]


def _tf(train: bool):
    aug = [transforms.RandomHorizontalFlip(),
           transforms.RandomRotation(7)] if train else []
    return transforms.Compose([
        transforms.Grayscale(num_output_channels=3),
        transforms.Resize((224, 224)),
        *aug,
        transforms.ToTensor(),
        transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD),
    ])


def make_loaders(root: str, batch_size: int = 32, num_workers: int = 4):
    train_ds = datasets.ImageFolder(f"{root}/train", _tf(True))
    val_ds = datasets.ImageFolder(f"{root}/val", _tf(False))
    test_ds = datasets.ImageFolder(f"{root}/test", _tf(False))
    kw = dict(batch_size=batch_size, num_workers=num_workers, pin_memory=True)
    return (
        DataLoader(train_ds, shuffle=True, **kw),
        DataLoader(val_ds, shuffle=False, **kw),
        DataLoader(test_ds, shuffle=False, **kw),
        train_ds.classes,
    )