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