| 9 | |
| 10 | |
| 11 | def 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 | |