(source_path, target_path, evaluation_path, transforms,
batch_size=32, return_id=False, balanced=False, val=False, val_data=None)
| 6 | |
| 7 | |
| 8 | def get_loader(source_path, target_path, evaluation_path, transforms, |
| 9 | batch_size=32, return_id=False, balanced=False, val=False, val_data=None): |
| 10 | source_folder = ImageFolder(os.path.join(source_path), |
| 11 | transforms[source_path], |
| 12 | return_id=return_id) |
| 13 | target_folder_train = ImageFolder(os.path.join(target_path), |
| 14 | transform=transforms[target_path], |
| 15 | return_paths=False, return_id=return_id) |
| 16 | if val: |
| 17 | source_val_train = ImageFolder(val_data, transforms[source_path], return_id=return_id) |
| 18 | target_folder_train = torch.utils.data.ConcatDataset([target_folder_train, source_val_train]) |
| 19 | source_val_test = ImageFolder(val_data, transforms[evaluation_path], return_id=return_id) |
| 20 | eval_folder_test = ImageFolder(os.path.join(evaluation_path), |
| 21 | transform=transforms["eval"], |
| 22 | return_paths=True) |
| 23 | |
| 24 | if balanced: |
| 25 | freq = Counter(source_folder.labels) |
| 26 | class_weight = {x: 1.0 / freq[x] for x in freq} |
| 27 | source_weights = [class_weight[x] for x in source_folder.labels] |
| 28 | sampler = WeightedRandomSampler(source_weights, |
| 29 | len(source_folder.labels)) |
| 30 | print("use balanced loader") |
| 31 | source_loader = torch.utils.data.DataLoader( |
| 32 | source_folder, |
| 33 | batch_size=batch_size, |
| 34 | sampler=sampler, |
| 35 | drop_last=True, |
| 36 | num_workers=4) |
| 37 | else: |
| 38 | source_loader = torch.utils.data.DataLoader( |
| 39 | source_folder, |
| 40 | batch_size=batch_size, |
| 41 | shuffle=True, |
| 42 | drop_last=True, |
| 43 | num_workers=4) |
| 44 | |
| 45 | target_loader = torch.utils.data.DataLoader( |
| 46 | target_folder_train, |
| 47 | batch_size=batch_size, |
| 48 | shuffle=True, |
| 49 | drop_last=True, |
| 50 | num_workers=4) |
| 51 | test_loader = torch.utils.data.DataLoader( |
| 52 | eval_folder_test, |
| 53 | batch_size=batch_size, |
| 54 | shuffle=False, |
| 55 | num_workers=4) |
| 56 | if val: |
| 57 | test_loader_source = torch.utils.data.DataLoader( |
| 58 | source_val_test, |
| 59 | batch_size=batch_size, |
| 60 | shuffle=False, |
| 61 | num_workers=4) |
| 62 | return source_loader, target_loader, test_loader, test_loader_source |
| 63 | |
| 64 | return source_loader, target_loader, test_loader, target_folder_train |
| 65 |
no test coverage detected