| 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: |
| |
| dataset = datasets.EuroSAT(root=data_dir, download=False, transform=EUROSAT_TRANSFORM) |
| except (RuntimeError, FileNotFoundError): |
| |
| 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]) |
|
|
| |
| 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 = [] |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| 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) |
| |
| |
| |
| return image, np.array(mask) |