ecopulse / src /data_loader.py
acibZ's picture
Deploy EcoPulse
43abac3
Raw
History Blame Contribute Delete
3.45 kB
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)