| 189 | |
| 190 | |
| 191 | def get_val_loader(engine, dataset, config, val_batch_size=1): |
| 192 | data_setting = { |
| 193 | "rgb_root": config.rgb_root_folder, |
| 194 | "rgb_format": config.rgb_format, |
| 195 | "gt_root": config.gt_root_folder, |
| 196 | "gt_format": config.gt_format, |
| 197 | "transform_gt": config.gt_transform, |
| 198 | "x_root": config.x_root_folder, |
| 199 | "x_format": config.x_format, |
| 200 | "x_single_channel": config.x_is_single_channel, |
| 201 | "class_names": config.class_names, |
| 202 | "train_source": config.train_source, |
| 203 | "eval_source": config.eval_source, |
| 204 | "class_names": config.class_names, |
| 205 | "dataset_name": config.dataset_name, |
| 206 | "backbone": config.backbone, |
| 207 | } |
| 208 | val_preprocess = ValPre(config.norm_mean, config.norm_std, config.x_is_single_channel, config) |
| 209 | |
| 210 | val_dataset = dataset(data_setting, "val", val_preprocess) |
| 211 | |
| 212 | val_sampler = None |
| 213 | is_shuffle = False |
| 214 | batch_size = val_batch_size |
| 215 | |
| 216 | if engine.distributed: |
| 217 | val_sampler = torch.utils.data.distributed.DistributedSampler(val_dataset) |
| 218 | batch_size = val_batch_size // engine.world_size |
| 219 | is_shuffle = False |
| 220 | |
| 221 | val_loader = data.DataLoader( |
| 222 | val_dataset, |
| 223 | batch_size=batch_size, |
| 224 | num_workers=config.num_workers, |
| 225 | drop_last=False, |
| 226 | shuffle=is_shuffle, |
| 227 | pin_memory=True, |
| 228 | sampler=val_sampler, |
| 229 | # worker_init_fn=seed_worker, |
| 230 | # generator=g, |
| 231 | ) |
| 232 | |
| 233 | return val_loader, val_sampler |