CoolFace
Apppublic

shiva-1993/transfer-learning-project

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
data.py215 linesDownload Raw Back to utils
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