File size: 3,449 Bytes
43abac3 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 | 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) |