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)