Spaces:
Running on Zero
Running on Zero
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
|