shiva-1993/transfer-learning-project
0
1"""2Shared data utilities: EuroSAT loading, stratified fraction sampling,3augmentation pipelines, DataLoader construction.4"""5 6from __future__ import annotations7 8import random9 10from datasets import load_dataset11from PIL import Image12from torch.utils.data import DataLoader, Dataset, Subset13from torchvision import transforms14 15# ── EuroSAT label map ─────────────────────────────────────────────────────────16 17EUROSAT_LABEL2ID = {18 "AnnualCrop": 0,19 "Forest": 1,20 "HerbaceousVegetation": 2,21 "Highway": 3,22 "Industrial": 4,23 "Pasture": 5,24 "PermanentCrop": 6,25 "Residential": 7,26 "River": 8,27 "SeaLake": 9,28}29EUROSAT_ID2LABEL = {v: k for k, v in EUROSAT_LABEL2ID.items()}30 31 32# ── Augmentation pipelines ────────────────────────────────────────────────────33 34 35def get_train_transform(36 image_size: int = 224, strength: str = "medium"37) -> transforms.Compose:38 """39 Training augmentation. Strength controls how aggressive the augmentation is.40 Satellite imagery benefits from rotation/flip but not colour jitter as strongly.41 """42 base = [43 transforms.Resize((image_size, image_size)),44 transforms.RandomHorizontalFlip(),45 transforms.RandomVerticalFlip(),46 ]47 48 if strength == "light":49 extra = [transforms.RandomRotation(10)]50 elif strength == "medium":51 extra = [52 transforms.RandomRotation(15),53 transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.1),54 transforms.RandomAffine(degrees=0, translate=(0.05, 0.05)),55 ]56 else: # strong57 extra = [58 transforms.RandomRotation(30),59 transforms.ColorJitter(60 brightness=0.3, contrast=0.3, saturation=0.2, hue=0.0561 ),62 transforms.RandomAffine(degrees=0, translate=(0.1, 0.1), scale=(0.9, 1.1)),63 transforms.GaussianBlur(kernel_size=3, sigma=(0.1, 1.0)),64 ]65 66 return transforms.Compose(67 base68 + extra69 + [70 transforms.ToTensor(),71 transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),72 ]73 )74 75 76def get_val_transform(image_size: int = 224) -> transforms.Compose:77 return transforms.Compose(78 [79 transforms.Resize((image_size, image_size)),80 transforms.ToTensor(),81 transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),82 ]83 )84 85 86# ── PyTorch Dataset wrapper ────────────────────────────────────────────────────87 88 89class EuroSATDataset(Dataset):90 """Wraps a HuggingFace EuroSAT split as a PyTorch Dataset."""91 92 def __init__(self, hf_split, transform=None):93 self.data = hf_split94 self.transform = transform95 96 def __len__(self):97 return len(self.data)98 99 def __getitem__(self, idx):100 item = self.data[idx]101 image = item["image"]102 if not isinstance(image, Image.Image):103 image = Image.fromarray(image).convert("RGB")104 else:105 image = image.convert("RGB")106 107 label = item["label"]108 109 if self.transform:110 image = self.transform(image)111 112 return image, label113 114 115# ── Data loading ──────────────────────────────────────────────────────────────116 117 118def load_eurosat(119 dataset_name: str = "timm/eurosat-rgb",120 data_fraction: float = 1.0,121 image_size: int = 224,122 batch_size: int = 32,123 num_workers: int = 4,124 augmentation_strength: str = "medium",125 seed: int = 42,126) -> tuple[DataLoader, DataLoader, DataLoader]:127 """128 Load EuroSAT, optionally subsample a stratified fraction of the training set,129 and return (train_loader, val_loader, test_loader).130 131 Args:132 data_fraction: Fraction of training data to use (1%, 5%, 10%, 100%).133 Validation and test sets are always full.134 135 Returns:136 (train_loader, val_loader, test_loader)137 """138 ds = load_dataset(dataset_name)139 140 train_ds = EuroSATDataset(141 ds["train"], transform=get_train_transform(image_size, augmentation_strength)142 )143 val_ds = EuroSATDataset(ds["validation"], transform=get_val_transform(image_size))144 test_ds = EuroSATDataset(ds["test"], transform=get_val_transform(image_size))145 146 # Stratified subsample of training set147 if data_fraction < 1.0:148 train_ds = _stratified_subset(train_ds, data_fraction, seed)149 150 train_loader = DataLoader(151 train_ds,152 batch_size=batch_size,153 shuffle=True,154 num_workers=num_workers,155 pin_memory=True,156 # Drop the ragged last batch only when there's more than one batch's157 # worth of data; otherwise a tiny-fraction subset would yield an empty158 # loader.159 drop_last=len(train_ds) > batch_size,160 )161 val_loader = DataLoader(162 val_ds,163 batch_size=batch_size,164 shuffle=False,165 num_workers=num_workers,166 pin_memory=True,167 )168 test_loader = DataLoader(169 test_ds,170 batch_size=batch_size,171 shuffle=False,172 num_workers=num_workers,173 pin_memory=True,174 )175 176 return train_loader, val_loader, test_loader177 178 179def _read_labels_fast(dataset) -> list[int]:180 """Read integer labels without decoding images.181 182 EuroSATDataset wraps a HuggingFace split whose ``label`` column can be read183 directly — accessing a single column does NOT decode the image column. The184 previous approach called ``dataset[i]`` for every item, which ran the full185 resize/augment/ToTensor pipeline on all ~16k images just to read their186 labels. Falls back to per-item access for plain datasets (e.g. test doubles)187 that don't expose the underlying column.188 """189 hf_split = getattr(dataset, "data", None)190 if hf_split is not None:191 try:192 return [int(x) for x in hf_split["label"]]193 except (KeyError, TypeError):194 pass195 labels = getattr(dataset, "labels", None)196 if labels is not None:197 return [int(x) for x in labels]198 return [int(dataset[i][1]) for i in range(len(dataset))]199 200 201def _stratified_subset(dataset: EuroSATDataset, fraction: float, seed: int) -> Subset:202 """Return a stratified subset keeping `fraction` of each class."""203 rng = random.Random(seed)204 labels = _read_labels_fast(dataset)205 class_indices: dict[int, list[int]] = {}206 for idx, lbl in enumerate(labels):207 class_indices.setdefault(lbl, []).append(idx)208 209 selected = []210 for cls_idxs in class_indices.values():211 n = max(1, int(len(cls_idxs) * fraction))212 selected.extend(rng.sample(cls_idxs, n))213 214 return Subset(dataset, selected)215 