File size: 9,747 Bytes
2407511
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
import json
import random
from enum import Enum
from pathlib import Path
from typing import Callable, List, Literal, Optional, Tuple, Union

import torch
from PIL import Image
from torch.utils.data import Dataset
from torchvision.transforms import v2


class SplitType(Enum):
    """Enumeration for dataset split types"""

    TRAIN = "train"
    VAL = "val"
    TEST = "test"


class Sentinel(Dataset):
    """
    A PyTorch Dataset for handling Sentinel-1&2 Image Pairs.

    This dataset assumes a directory structure of:
    root_dir/
        category1/
            s1/
                image1.png
                image2.png
            s2/
                image1.png
                image2.png
        category2/
            ...

    This class has support for train/val/test splits. When `split_type` is `None`,
    uses the complete dataset. When `split_type` is specified
    (``'train'``, ``'val'``, ``'test'``), the dataset can be split using:

    1. A split that defines which images belong to which split
    2. Random splitting with a specified ratio

    Args:
        root_dir (str | Path): Root directory containing the dataset
        split_type (str | None): Which split to use ('train', 'val', 'test') or None for full dataset
        transform (callable, optional): Transform to apply to both SAR and optical images
        split_mode (str, optional): How to split the dataset ('random', 'split')
        split_ratio (Tuple[float, float, float], optional): Ratio for train/val/test splits
        split_file (str | Path, optional): predefined the splits
        seed (int, optional): Random seed for reproducible splitting

    Attributes:
        root_dir (Path): Path to the dataset root directory
        transform (callable): Transform pipeline for the images
        image_pairs (List[Tuple[Path, Path]]): List of paired image paths (SAR, optical)
    """

    def __init__(
        self,
        root_dir: Union[str, Path],
        split_type: Optional[str] = None,
        input_transform: Optional[Callable] = None,
        target_transform: Optional[Callable] = None,
        split_mode: Literal["random", "split"] = "random",
        split_ratio: Tuple[float, float, float] = (0.7, 0.15, 0.15),
        split_file: Optional[Union[str, Path]] = None,
        seed: int = 42,
    ):
        self.root_dir = Path(root_dir)
        if not self.root_dir.exists():
            raise FileNotFoundError(
                f"Dataset root directory not found: {self.root_dir}"
            )

        # Convert string split_type to enum if provided
        self.split_type = SplitType(split_type) if split_type else None

        # Default transform pipeline
        self.input_transform = (
            input_transform
            if input_transform
            else v2.Compose([v2.ToImage(), v2.ToDtype(torch.float32, scale=True)])
        )
        self.target_transform = (
            target_transform if target_transform else self.input_transform
        )

        # Collect image pairs
        self.all_image_pairs = self._collect_images()

        # Apply split if specified
        if split_type:
            if split_mode == "split" and split_file:
                self.image_pairs = self._apply_predefined_split(split_file)
            elif split_mode == "random":
                self.image_pairs = self._apply_random_split(split_ratio, seed)
            else:
                raise ValueError(
                    "Invalid split configuration. Use either 'split' with a split_file or 'random' with split_ratio"
                )
        else:
            # If no split type specified, use all images
            self.image_pairs = self.all_image_pairs

        print(f"Total image pairs found: {len(self)}")

    def _collect_images(self) -> List[Tuple[Path, Path]]:
        """
        Collects paired SAR (s1) and optical (s2) image paths from the dataset directory.

        Returns:
            List[Tuple[Path, Path]]: List of (SAR image path, optical image path) pairs
        """
        image_pairs = []

        # Iterate through category subdirectories
        for category in self.root_dir.iterdir():
            # Check if it's a directory
            if not category.is_dir():
                continue

            s1_path = category / "s1"
            s2_path = category / "s2"

            if not (s1_path.is_dir() and s2_path.is_dir()):
                # print(f"Missing s1 or s2 subdirectory in category: {category.name}")
                continue

            # Collect pairs
            for s1_file in s1_path.glob("*.png"):
                # Convert SAR filename to optical filename
                # e.g. 'ROIs1970_fall_s1_13_p265.png' -> 'ROIs1970_fall_s2_13_p265.png'
                s2_filename = list(s1_file.name.split("_"))
                s2_filename[2] = "s2"
                s2_file = s2_path / "_".join(s2_filename)

                if not s2_file.exists():
                    # print(f"Missing optical image for SAR image: {s1_file.name} - {s2_file.name}")
                    continue

                image_pairs.append((s1_file, s2_file))

        return image_pairs

    def _apply_predefined_split(
        self, split_file: Union[str, Path]
    ) -> List[Tuple[Path, Path]]:
        """
        Applies a predefined split from a JSON file.

        Args:
            split_file: Path to JSON file containing split definitions

        Returns:
            List[Tuple[Path, Path]]: Image pairs for the specified split
        """
        try:
            with open(split_file, "r") as f:  # get the split content
                splits = json.load(f)

            if self.split_type.value not in splits["data"]:  # check if it helds
                raise ValueError(
                    f"Split type {self.split_type.value} not found in split file"
                )

            split_filenames = set(
                splits["data"][self.split_type.value]
            )  # data['split']
            return [
                pair
                for pair in self.all_image_pairs  # collect and return split
                if any(
                    str(p.relative_to(self.root_dir)) in split_filenames
                    for p in pair[:2]
                )
            ]
        except Exception as e:
            print(f"Could not open split file\n\t{e}")
            raise

    def _apply_random_split(
        self, split_ratio: Tuple[float, float, float], seed: int
    ) -> List[Tuple[Path, Path]]:
        """
        Randomly splits the dataset according to the given ratios.

        Args:
            split_ratio: Tuple of (train, val, test) ratios
            seed: Random seed for reproducibility

        Returns:
            List[Tuple[Path, Path]]: Image pairs for the specified split
        """
        if sum(split_ratio) != 1:
            raise ValueError("Split ratios must sum to 1")

        # Set random seed for reproducibility
        random.seed(seed)

        # Shuffle indices
        indices = list(range(len(self.all_image_pairs)))
        random.shuffle(indices)

        # Calculate split points
        train_end = int(len(indices) * split_ratio[0])
        val_end = train_end + int(len(indices) * split_ratio[1])

        # Select appropriate slice based on split type
        if self.split_type == SplitType.TRAIN:
            split_indices = indices[:train_end]
        elif self.split_type == SplitType.VAL:
            split_indices = indices[train_end:val_end]
        else:  # TEST
            split_indices = indices[val_end:]

        return [self.all_image_pairs[i] for i in split_indices]

    def save_split(self, output_file: Union[str, Path], is_exists: bool = False):
        """
        Saves the current split configuration to a JSON file.

        Args:
            output_file: Path to save the split configuration
            is_exists: If file exist, add new split data
        """
        if self.split_type:
            split = self.split_type.value
            split_info = {
                "data": {
                    split: [
                        str(p[0].relative_to(self.root_dir)) for p in self.image_pairs
                    ]
                }
            }
            # Check if the file exists
            if is_exists and Path(output_file).exists():
                # Read the existing content
                with open(output_file, "r") as f:
                    existing_data = json.load(f)
                # Check if 'data' is already in the existing content, if not, create it
                if "data" not in existing_data:
                    existing_data["data"] = {}

                # Add or update the split information
                existing_data["data"][split] = split_info["data"][split]
                split_info = existing_data

            with open(output_file, "w") as f:
                json.dump(split_info, f, indent=2)

    def __len__(self):
        """Returns the total number of image pairs in the dataset."""
        return len(self.image_pairs)

    def __getitem__(self, idx: int) -> Tuple[torch.Tensor, torch.Tensor]:
        """
        Retrieves the image pair at the given index.

        Args:
            idx (int): Index of the image pair to retrieve

        Returns:
            Tuple[torch.Tensor, torch.Tensor]: Processed (SAR image, optical image) pair
        """
        # Get paths for SAR and optical images
        s1_path, s2_path = self.image_pairs[idx]

        # Load images
        s1_image = Image.open(s1_path).convert("RGB")
        s2_image = Image.open(s2_path).convert("RGB")

        # Apply transforms
        s1_image = self.input_transform(s1_image)
        s2_image = self.target_transform(s2_image)

        return s1_image, s2_image