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

Function get_dataloaders

utils/defaults.py:11–51  ·  view source on GitHub ↗
(kwargs)

Source from the content-addressed store, hash-verified

9
10
11def get_dataloaders(kwargs):
12 source_data = kwargs["source_data"]
13 target_data = kwargs["target_data"]
14 evaluation_data = kwargs["evaluation_data"]
15 conf = kwargs["conf"]
16 val_data = None
17 if "val" in kwargs:
18 val = kwargs["val"]
19 if val:
20 val_data = kwargs["val_data"]
21 else:
22 val = False
23
24 data_transforms = {
25 source_data: transforms.Compose([
26 transforms.Scale((256, 256)),
27 transforms.RandomHorizontalFlip(),
28 transforms.RandomCrop(224),
29 transforms.ToTensor(),
30 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
31 ]),
32 target_data: transforms.Compose([
33 transforms.Scale((256, 256)),
34 transforms.RandomHorizontalFlip(),
35 transforms.RandomCrop(224),
36 transforms.ToTensor(),
37 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
38 ]),
39 "eval": transforms.Compose([
40 transforms.Scale((256, 256)),
41 transforms.CenterCrop(224),
42 transforms.ToTensor(),
43 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
44 ]),
45 }
46 return get_loader(source_data, target_data, evaluation_data,
47 data_transforms,
48 batch_size=conf.data.dataloader.batch_size,
49 return_id=True,
50 balanced=conf.data.dataloader.class_balance,
51 val=val, val_data=val_data)
52
53
54

Callers 1

train.pyFile · 0.90

Calls 1

get_loaderFunction · 0.90

Tested by

no test coverage detected