diff options
Diffstat (limited to 'dataset.py')
| -rw-r--r-- | dataset.py | 70 |
1 files changed, 51 insertions, 19 deletions
| @@ -4,7 +4,7 @@ import numpy as np | |||
| 4 | import cv2 | 4 | import cv2 |
| 5 | 5 | ||
| 6 | import torch | 6 | import torch |
| 7 | from torch.utils.data import Dataset, DataLoader | 7 | from torch.utils.data import Dataset, DataLoader, WeightedRandomSampler |
| 8 | from torchvision import transforms | 8 | from torchvision import transforms |
| 9 | from sklearn.model_selection import train_test_split | 9 | from 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 | ||
| 80 | def build_transforms(use_clahe: bool, training: bool): | 78 | def 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 |
