File size: 6,412 Bytes
37fbec9 | 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 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 | """
data/dataset.py
---------------
PyTorch Dataset for NIH ChestX-ray14.
Key design decisions:
- Uses Hugging Face `datasets` for streaming from `BahaaEldin0/NIH-Chest-Xray-14`.
- Avoids 45GB local download, scales natively up to 112k images.
- Multi-hot label encoding for 14 simultaneous disease labels.
- CLAHE preprocessing baked into the transform pipeline.
"""
import torch
from torch.utils.data import IterableDataset, DataLoader
from datasets import load_dataset
import numpy as np
import logging
import cv2
from PIL import Image
from data.preprocessing import get_train_transforms, get_val_transforms, apply_clahe
logger = logging.getLogger(__name__)
# ββ 14 NIH disease labels βββββββββββββββββββββββββββββββββββββββββββββββββββββ
DISEASE_LABELS = [
"Atelectasis",
"Cardiomegaly",
"Effusion",
"Infiltration",
"Mass",
"Nodule",
"Pneumonia",
"Pneumothorax",
"Consolidation",
"Edema",
"Emphysema",
"Fibrosis",
"Pleural_Thickening",
"Hernia",
]
NUM_CLASSES = len(DISEASE_LABELS)
def encode_labels(labels) -> np.ndarray:
"""
Convert a list of string labels (or a pipe separated string) into a 14-dim multi-hot vector.
"""
if isinstance(labels, str):
if labels == "No Finding":
labels_list = []
else:
labels_list = labels.split("|")
else:
labels_list = labels
label_vec = np.zeros(NUM_CLASSES, dtype=np.float32)
for disease in labels_list:
disease = disease.strip()
if disease in DISEASE_LABELS:
label_vec[DISEASE_LABELS.index(disease)] = 1.0
return label_vec
class HFStreamingChestXrayDataset(IterableDataset):
"""
Wrapper around Hugging Face IterableDataset for PyTorch DataLoader compatibility.
"""
def __init__(
self,
hf_iterable,
image_size: int = 224,
transform=None,
clip_limit: float = 2.0,
tile_size: int = 8,
take_limit: int = None
):
self.hf_iterable = hf_iterable
self.image_size = image_size
self.transform = transform
self.clip_limit = clip_limit
self.tile_size = tile_size
self.take_limit = take_limit
def __iter__(self):
iterable = self.hf_iterable
if self.take_limit is not None:
iterable = iterable.take(self.take_limit)
for idx, item in enumerate(iterable):
image = item.get('image') or item.get('Image')
if image is None:
continue
# Ensure image is PIL Image and convert to grayscale numpy array
if not isinstance(image, Image.Image):
continue
img_gray = np.array(image.convert("L"))
# CLAHE
img_clahe = apply_clahe(img_gray, clip_limit=self.clip_limit, tile_size=self.tile_size)
# Resize BEFORE stacking to save memory/computation
img_resized = cv2.resize(img_clahe, (self.image_size, self.image_size), interpolation=cv2.INTER_AREA)
# Convert to RGB numpy array for Albumentations
img_rgb = np.stack([img_resized, img_resized, img_resized], axis=-1)
# Transform
if self.transform:
augmented = self.transform(image=img_rgb)
img_tensor = augmented["image"]
else:
img_tensor = torch.from_numpy(img_rgb.transpose(2, 0, 1)).float() / 255.0
# Labels
label_str_list = item.get('label', [])
label_vec = encode_labels(label_str_list)
label_tensor = torch.tensor(label_vec, dtype=torch.float32)
# Yield path as Patient ID or fallback to str index to satisfy interface
patient_id = str(item.get('Patient ID', idx))
yield img_tensor, label_tensor, patient_id
def build_dataloaders(
images_dir=None,
labels_csv=None,
train_list_txt=None,
test_list_txt=None,
image_size: int = 224,
batch_size: int = 8,
num_workers: int = 0,
pin_memory: bool = True,
val_split: float = 0.1,
train_fraction: float = 1.0,
seed: int = 42,
train_take_limit: int = 20000,
val_take_limit: int = 2000,
):
"""
Build train, validation, and test DataLoaders directly from the HF Stream.
Uses generic fallback args for compatibility with old code calls.
"""
print(f" [Dataset] Loading HF Stream BahaaEldin0/NIH-Chest-Xray-14...")
# In order to stream properly we load each split
ds = load_dataset('BahaaEldin0/NIH-Chest-Xray-14', streaming=True)
train_transform = get_train_transforms(image_size)
val_transform = get_val_transforms(image_size)
# Note: BahaaEldin0 dataset has 'train', 'valid', 'test' splits
train_stream = ds['train']
val_stream = ds['valid']
test_stream = ds['test']
# Shuffle the training stream slightly for better randomness (buffer_size=1000)
train_stream = train_stream.shuffle(buffer_size=1000, seed=seed)
# If train_fraction < 1.0, we can adjust take limits accordingly.
if train_fraction < 1.0 and train_take_limit is not None:
train_take_limit = int(train_take_limit * train_fraction)
train_ds = HFStreamingChestXrayDataset(train_stream, image_size=image_size, transform=train_transform, take_limit=train_take_limit)
val_ds = HFStreamingChestXrayDataset(val_stream, image_size=image_size, transform=val_transform, take_limit=val_take_limit)
test_ds = HFStreamingChestXrayDataset(test_stream, image_size=image_size, transform=val_transform, take_limit=val_take_limit)
dl_kwargs = dict(
batch_size=batch_size,
num_workers=num_workers,
pin_memory=pin_memory,
)
train_loader = DataLoader(train_ds, **dl_kwargs)
val_loader = DataLoader(val_ds, **dl_kwargs)
test_loader = DataLoader(test_ds, **dl_kwargs)
# Calculate uniform weights (streaming mode makes exact calculation slow)
class_weights = torch.ones(NUM_CLASSES, dtype=torch.float32)
print(f"\n DataLoaders ready:")
print(f" Train limit: {train_take_limit if train_take_limit else 'Full'}")
print(f" Val/Test limit: {val_take_limit if val_take_limit else 'Full'}")
return train_loader, val_loader, test_loader, class_weights
|