| 178 | |
| 179 | |
| 180 | class DynamicEarthNet(SitsDataset): |
| 181 | def __init__( |
| 182 | self, |
| 183 | path, |
| 184 | split="train", |
| 185 | domain_shift=False, |
| 186 | num_channels=4, |
| 187 | num_classes=7, |
| 188 | img_size=128, |
| 189 | true_size=1024, |
| 190 | train_length=6, |
| 191 | date_aug_range=2 |
| 192 | ): |
| 193 | |
| 194 | super(DynamicEarthNet, self).__init__(path=path, |
| 195 | split=split, |
| 196 | domain_shift=domain_shift, |
| 197 | num_channels=num_channels, |
| 198 | num_classes=num_classes, |
| 199 | img_size=img_size, |
| 200 | true_size=true_size, |
| 201 | train_length=train_length) |
| 202 | """Initializes the dataset. |
| 203 | Args: |
| 204 | path (str): path to the dataset |
| 205 | split (str): split to use (train, val, test) |
| 206 | domain_shift (bool): if val/test, whether we are in a domain shift setting or not |
| 207 | """ |
| 208 | self.date_aug_range = date_aug_range |
| 209 | self.monthly_dates = get_monthly_dates_dict() |
| 210 | self.gt, self.sits_ids = self.load_ground_truth(split) |
| 211 | self.month_list = [list(range(24)) for _ in range(self.gt.shape[0])] |
| 212 | self.mean = torch.tensor([83.1029, 80.7615, 69.3328, 133.8648], dtype=torch.float16).reshape(4, 1, 1) |
| 213 | self.std = torch.tensor([33.2714, 25.5288, 23.9868, 30.5591], dtype=torch.float16).reshape(4, 1, 1) |
| 214 | |
| 215 | def load_data(self, sits_number, sits_id, months, curr_sits_path): |
| 216 | data = torch.zeros((len(months), self.num_channels, self.true_size, self.true_size), dtype=torch.float16) |
| 217 | days = [self.random_date_augmentation(month) for month in months] |
| 218 | name_rgb = [f'{sits_id}_{day}_rgb.jpeg' for day in days] |
| 219 | name_infra = [f'{sits_id}_{day}_infra.jpeg' for day in days] |
| 220 | for d, (n_rgb, n_infra) in enumerate(zip(name_rgb, name_infra)): |
| 221 | data[d, :3] = torchvision.io.read_image(join(curr_sits_path, n_rgb)) |
| 222 | data[d, 3] = torchvision.io.read_image(join(curr_sits_path, n_infra)) |
| 223 | return data, days |
| 224 | |
| 225 | def random_date_augmentation(self, month): |
| 226 | if self.split == 'train': |
| 227 | return max(0, random.randint(0, self.date_aug_range * 2) - self.date_aug_range + self.monthly_dates[month]) |
| 228 | else: |
| 229 | return self.monthly_dates[month] |
| 230 | |
| 231 | |
| 232 | def collate_fn(batch): |
nothing calls this directly
no outgoing calls
no test coverage detected