(config, args)
| 67 | return dataset_val, data_loader_val |
| 68 | |
| 69 | def build_loader(config, args): |
| 70 | config.defrost() |
| 71 | dataset_train, config.MODEL.NUM_CLASSES = build_dataset(is_train=True, config=config) |
| 72 | config.freeze() |
| 73 | # print(f"local rank {config.LOCAL_RANK} / global rank {dist.get_rank()} successfully build train dataset") |
| 74 | dataset_val, _ = build_dataset(is_train=False, config=config) |
| 75 | # print(f"local rank {config.LOCAL_RANK} / global rank {dist.get_rank()} successfully build val dataset") |
| 76 | |
| 77 | num_tasks = dist.get_world_size() |
| 78 | global_rank = dist.get_rank() |
| 79 | |
| 80 | sampler_train = torch.utils.data.DistributedSampler( |
| 81 | dataset_train, num_replicas=num_tasks, rank=global_rank, shuffle=True |
| 82 | ) |
| 83 | |
| 84 | if config.TEST.SEQUENTIAL: |
| 85 | sampler_val = torch.utils.data.SequentialSampler(dataset_val) |
| 86 | else: |
| 87 | sampler_val = torch.utils.data.distributed.DistributedSampler( |
| 88 | dataset_val, shuffle=False |
| 89 | ) |
| 90 | |
| 91 | data_loader_train = torch.utils.data.DataLoader( |
| 92 | dataset_train, sampler=sampler_train, |
| 93 | batch_size=args.calib_batchsz if args.calib else config.DATA.BATCH_SIZE, |
| 94 | num_workers=config.DATA.NUM_WORKERS, |
| 95 | pin_memory=config.DATA.PIN_MEMORY, |
| 96 | drop_last=True, |
| 97 | ) |
| 98 | |
| 99 | data_loader_val = torch.utils.data.DataLoader( |
| 100 | dataset_val, sampler=sampler_val, |
| 101 | batch_size=config.DATA.BATCH_SIZE, |
| 102 | shuffle=False, |
| 103 | num_workers=config.DATA.NUM_WORKERS, |
| 104 | pin_memory=config.DATA.PIN_MEMORY, |
| 105 | drop_last=True |
| 106 | ) |
| 107 | |
| 108 | return dataset_train, dataset_val, data_loader_train, data_loader_val |
| 109 | |
| 110 | |
| 111 | def build_dataset(is_train, config): |
no test coverage detected