(config)
| 19 | from .samplers import SubsetRandomSampler |
| 20 | from torch.utils.data import DataLoader |
| 21 | def build_loader(config): |
| 22 | |
| 23 | config.defrost() |
| 24 | dataset_train, config.MODEL.NUM_CLASSES = build_dataset(is_train=True, config=config) |
| 25 | config.freeze() |
| 26 | |
| 27 | dataset_val, _ = build_dataset(is_train=False, config=config) |
| 28 | |
| 29 | num_tasks = dist.get_world_size() |
| 30 | global_rank = dist.get_rank() |
| 31 | sampler_train = torch.utils.data.DistributedSampler( |
| 32 | dataset_train, num_replicas=num_tasks, rank=global_rank, shuffle=True |
| 33 | ) |
| 34 | indices = np.arange(dist.get_rank(), len(dataset_val), dist.get_world_size()) |
| 35 | sampler_val = SubsetRandomSampler(indices) |
| 36 | |
| 37 | data_loader_train = DataLoader( |
| 38 | dataset_train, sampler=sampler_train, |
| 39 | batch_size=config.DATA.BATCH_SIZE, |
| 40 | num_workers=config.DATA.NUM_WORKERS, |
| 41 | pin_memory=config.DATA.PIN_MEMORY, |
| 42 | drop_last=True |
| 43 | ) |
| 44 | |
| 45 | data_loader_val = DataLoader( |
| 46 | dataset_val, sampler=sampler_val, |
| 47 | batch_size=config.DATA.BATCH_SIZE, |
| 48 | shuffle=False, |
| 49 | num_workers=config.DATA.NUM_WORKERS, |
| 50 | pin_memory=config.DATA.PIN_MEMORY, |
| 51 | drop_last=False |
| 52 | ) |
| 53 | |
| 54 | # setup mixup / cutmix |
| 55 | mixup_fn = None |
| 56 | mixup_active = config.AUG.MIXUP > 0 or config.AUG.CUTMIX > 0. or config.AUG.CUTMIX_MINMAX is not None |
| 57 | if mixup_active: |
| 58 | mixup_fn = Mixup( |
| 59 | mixup_alpha=config.AUG.MIXUP, cutmix_alpha=config.AUG.CUTMIX, cutmix_minmax=config.AUG.CUTMIX_MINMAX, |
| 60 | prob=config.AUG.MIXUP_PROB, switch_prob=config.AUG.MIXUP_SWITCH_PROB, mode=config.AUG.MIXUP_MODE, |
| 61 | label_smoothing=config.MODEL.LABEL_SMOOTHING, num_classes=config.MODEL.NUM_CLASSES) |
| 62 | |
| 63 | return dataset_train, dataset_val, data_loader_train, data_loader_val, mixup_fn |
| 64 | |
| 65 | |
| 66 | def build_dataset(is_train, config): |
no test coverage detected