summaryrefslogtreecommitdiff
path: root/dataset.py
diff options
context:
space:
mode:
Diffstat (limited to 'dataset.py')
-rw-r--r--dataset.py70
1 files changed, 51 insertions, 19 deletions
diff --git a/dataset.py b/dataset.py
index 696a312..0068a97 100644
--- a/dataset.py
+++ b/dataset.py
@@ -4,7 +4,7 @@ import numpy as np
4import cv2 4import cv2
5 5
6import torch 6import torch
7from torch.utils.data import Dataset, DataLoader 7from torch.utils.data import Dataset, DataLoader, WeightedRandomSampler
8from torchvision import transforms 8from torchvision import transforms
9from sklearn.model_selection import train_test_split 9from sklearn.model_selection import train_test_split
10 10
@@ -19,14 +19,12 @@ class CLAHETransform:
19 """ 19 """
20 20
21 def __init__(self, clip_limit: float = 2.0, tile_grid_size: tuple = (8, 8)): 21 def __init__(self, clip_limit: float = 2.0, tile_grid_size: tuple = (8, 8)):
22 self.clip_limit = clip_limit 22 self.clahe = cv2.createCLAHE(clipLimit=clip_limit, tileGridSize=tile_grid_size)
23 self.tile_grid_size = tile_grid_size
24 23
25 def __call__(self, img: Image.Image) -> Image.Image: 24 def __call__(self, img: Image.Image) -> Image.Image:
26 img_np = np.array(img) # RGB uint8 25 img_np = np.array(img) # RGB uint8
27 lab = cv2.cvtColor(img_np, cv2.COLOR_RGB2LAB) 26 lab = cv2.cvtColor(img_np, cv2.COLOR_RGB2LAB)
28 clahe = cv2.createCLAHE(clipLimit=self.clip_limit, tileGridSize=self.tile_grid_size) 27 lab[:, :, 0] = self.clahe.apply(lab[:, :, 0])
29 lab[:, :, 0] = clahe.apply(lab[:, :, 0])
30 result = cv2.cvtColor(lab, cv2.COLOR_LAB2RGB) 28 result = cv2.cvtColor(lab, cv2.COLOR_LAB2RGB)
31 return Image.fromarray(result) 29 return Image.fromarray(result)
32 30
@@ -77,17 +75,32 @@ def split_dataset(data_root: str, seed: int = 42):
77 return train_paths, train_labels, val_paths, val_labels, test_paths, test_labels 75 return train_paths, train_labels, val_paths, val_labels, test_paths, test_labels
78 76
79 77
80def build_transforms(use_clahe: bool, training: bool): 78def build_transforms(use_clahe: bool, training: bool, img_size: int = 224):
81 """Build a transform pipeline for training or eval.""" 79 """Build a transform pipeline for training or eval.
82 ops = [transforms.Resize((224, 224))] 80
83 if use_clahe: 81 Training path uses RandomResizedCrop for scale/position augmentation.
84 ops.append(CLAHETransform(clip_limit=2.0, tile_grid_size=(8, 8))) 82 Eval path uses a deterministic Resize to preserve comparability.
83 CLAHE is applied before spatial transforms at a slightly larger size
84 so the crop has room to operate.
85 """
86 ops = []
85 if training: 87 if training:
88 if use_clahe:
89 # Resize larger so RandomResizedCrop still sees enough context after CLAHE
90 ops.append(transforms.Resize((img_size + 32, img_size + 32)))
91 ops.append(CLAHETransform(clip_limit=2.0, tile_grid_size=(8, 8)))
86 ops += [ 92 ops += [
93 transforms.RandomResizedCrop(img_size, scale=(0.8, 1.0), ratio=(0.9, 1.1)),
87 transforms.RandomHorizontalFlip(p=0.5), 94 transforms.RandomHorizontalFlip(p=0.5),
88 transforms.RandomVerticalFlip(p=0.5), 95 transforms.RandomVerticalFlip(p=0.5),
89 transforms.RandomRotation(degrees=15), 96 transforms.RandomRotation(degrees=30),
97 transforms.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.02),
98 transforms.GaussianBlur(kernel_size=3, sigma=(0.1, 1.0)),
90 ] 99 ]
100 else:
101 ops.append(transforms.Resize((img_size, img_size)))
102 if use_clahe:
103 ops.append(CLAHETransform(clip_limit=2.0, tile_grid_size=(8, 8)))
91 ops += [ 104 ops += [
92 transforms.ToTensor(), 105 transforms.ToTensor(),
93 transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD), 106 transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),
@@ -102,22 +115,41 @@ def build_dataloaders(
102 use_clahe: bool, 115 use_clahe: bool,
103 batch_size: int = 32, 116 batch_size: int = 32,
104 num_workers: int = 4, 117 num_workers: int = 4,
118 img_size: int = 224,
105): 119):
106 """Create DataLoaders for all three splits.""" 120 """Create DataLoaders for all three splits.
107 train_ds = DRDataset(train_paths, train_labels, transform=build_transforms(use_clahe, training=True)) 121
108 val_ds = DRDataset(val_paths, val_labels, transform=build_transforms(use_clahe, training=False)) 122 The training DataLoader uses WeightedRandomSampler so that each class
109 test_ds = DRDataset(test_paths, test_labels, transform=build_transforms(use_clahe, training=False)) 123 is sampled at approximately equal frequency, counteracting class imbalance.
124 """
125 train_ds = DRDataset(train_paths, train_labels, transform=build_transforms(use_clahe, training=True, img_size=img_size))
126 val_ds = DRDataset(val_paths, val_labels, transform=build_transforms(use_clahe, training=False, img_size=img_size))
127 test_ds = DRDataset(test_paths, test_labels, transform=build_transforms(use_clahe, training=False, img_size=img_size))
128
129 # Per-sample weights: inverse of class frequency so all classes are seen equally
130 labels_arr = np.array(train_labels)
131 class_counts = np.bincount(labels_arr, minlength=5)
132 class_sample_weights = 1.0 / class_counts.astype(float)
133 sample_weights = class_sample_weights[labels_arr]
134 sampler = WeightedRandomSampler(
135 weights=torch.from_numpy(sample_weights).float(),
136 num_samples=len(train_labels),
137 replacement=True,
138 )
110 139
111 train_loader = DataLoader( 140 train_loader = DataLoader(
112 train_ds, batch_size=batch_size, shuffle=True, 141 train_ds, batch_size=batch_size, sampler=sampler,
113 num_workers=num_workers, pin_memory=True 142 num_workers=num_workers, pin_memory=True,
143 persistent_workers=True, prefetch_factor=4,
114 ) 144 )
115 val_loader = DataLoader( 145 val_loader = DataLoader(
116 val_ds, batch_size=batch_size * 2, shuffle=False, 146 val_ds, batch_size=batch_size * 2, shuffle=False,
117 num_workers=num_workers, pin_memory=True 147 num_workers=num_workers, pin_memory=True,
148 persistent_workers=True, prefetch_factor=4,
118 ) 149 )
119 test_loader = DataLoader( 150 test_loader = DataLoader(
120 test_ds, batch_size=batch_size * 2, shuffle=False, 151 test_ds, batch_size=batch_size * 2, shuffle=False,
121 num_workers=num_workers, pin_memory=True 152 num_workers=num_workers, pin_memory=True,
153 persistent_workers=True, prefetch_factor=4,
122 ) 154 )
123 return train_loader, val_loader, test_loader 155 return train_loader, val_loader, test_loader