import os import torch from torch.utils.data import DataLoader, Dataset from torchvision import datasets from PIL import Image import numpy as np from src.transforms import EUROSAT_TRANSFORM def get_eurosat_dataloaders(data_dir, batch_size=32, train_split=0.8): """ Returns train and validation dataloaders for the EuroSAT dataset. """ try: # If the root has a 'eurosat' folder created by torchvision dataset = datasets.EuroSAT(root=data_dir, download=False, transform=EUROSAT_TRANSFORM) except (RuntimeError, FileNotFoundError): # Fallback to ImageFolder if extracted differently or torchvision layout differs dataset_path = os.path.join(data_dir, '2750') if not os.path.exists(dataset_path): dataset_path = data_dir dataset = datasets.ImageFolder(root=dataset_path, transform=transform) train_size = int(train_split * len(dataset)) val_size = len(dataset) - train_size train_dataset, val_dataset = torch.utils.data.random_split(dataset, [train_size, val_size]) # Colab optimization: Use more workers and pin_memory for faster CPU-to-GPU transfer num_workers = min(4, os.cpu_count() or 2) train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=num_workers, pin_memory=True) val_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=num_workers, pin_memory=True) return train_loader, val_loader, dataset.classes class DeepGlobeDataset(Dataset): """ Custom Dataset for DeepGlobe Land Cover Classification. Expects directories: images/ and masks/ """ def __init__(self, root_dir, transform=None): self.root_dir = root_dir self.images_dir = os.path.join(root_dir, 'images') self.masks_dir = os.path.join(root_dir, 'masks') self.transform = transform if os.path.exists(self.images_dir): all_images = [f for f in os.listdir(self.images_dir) if f.endswith('.jpg')] self.image_files = [] for img in all_images: mask_name = img.replace('sat.jpg', 'mask.png') if os.path.exists(os.path.join(self.masks_dir, mask_name)): self.image_files.append(img) else: self.image_files = [] # DeepGlobe Color Map (RGB to Class) # Urban: (0, 255, 255) # Agriculture: (255, 255, 0) # Rangeland/Pasture: (255, 0, 255) # Forest: (0, 255, 0) # Water: (0, 0, 255) # Barren: (255, 255, 255) # Unknown: (0, 0, 0) self.classes = [ "Urban", "Agriculture", "Rangeland", "Forest", "Water", "Barren", "Unknown" ] def __len__(self): return len(self.image_files) def __getitem__(self, idx): img_name = self.image_files[idx] img_path = os.path.join(self.images_dir, img_name) mask_name = img_name.replace('sat.jpg', 'mask.png') mask_path = os.path.join(self.masks_dir, mask_name) image = Image.open(img_path).convert("RGB") mask = Image.open(mask_path).convert("RGB") if self.transform: image = self.transform(image) # Mask would need separate handling if applying spatial transforms, # but for simple inference, returning numpy arrays is often sufficient. return image, np.array(mask)