(source_data, target_data, target_data_label, evaluation_data, conf)
| 53 | |
| 54 | |
| 55 | def get_dataloaders_label(source_data, target_data, target_data_label, evaluation_data, conf): |
| 56 | |
| 57 | data_transforms = { |
| 58 | source_data: transforms.Compose([ |
| 59 | transforms.Scale((256, 256)), |
| 60 | transforms.RandomHorizontalFlip(), |
| 61 | transforms.RandomCrop(224), |
| 62 | transforms.ToTensor(), |
| 63 | transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) |
| 64 | ]), |
| 65 | target_data: transforms.Compose([ |
| 66 | transforms.Scale((256, 256)), |
| 67 | transforms.RandomHorizontalFlip(), |
| 68 | transforms.RandomCrop(224), |
| 69 | transforms.ToTensor(), |
| 70 | transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) |
| 71 | ]), |
| 72 | evaluation_data: transforms.Compose([ |
| 73 | transforms.Scale((256, 256)), |
| 74 | transforms.CenterCrop(224), |
| 75 | transforms.ToTensor(), |
| 76 | transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) |
| 77 | ]), |
| 78 | } |
| 79 | return get_loader_label(source_data, target_data, target_data_label, |
| 80 | evaluation_data, data_transforms, |
| 81 | batch_size=conf.data.dataloader.batch_size, |
| 82 | return_id=True, |
| 83 | balanced=conf.data.dataloader.class_balance) |
| 84 | |
| 85 | def get_models(kwargs): |
| 86 | net = kwargs["network"] |
nothing calls this directly
no test coverage detected