MCPcopy Create free account
hub / github.com/VisionLearningGroup/OVANet / get_loader

Function get_loader

data_loader/get_loader.py:8–64  ·  view source on GitHub ↗
(source_path, target_path, evaluation_path, transforms,
               batch_size=32, return_id=False, balanced=False, val=False, val_data=None)

Source from the content-addressed store, hash-verified

6
7
8def 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

Callers 1

get_dataloadersFunction · 0.90

Calls 1

ImageFolderClass · 0.85

Tested by

no test coverage detected