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