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