| 70 | |
| 71 | |
| 72 | class Augment_SPH_Data(): |
| 73 | def __init__(self, img_size, mode=0, is_train=False): |
| 74 | if is_train: |
| 75 | if mode == 0: |
| 76 | transform = [transforms.ToTensor(), |
| 77 | transforms.Normalize([0.5], [0.5])] |
| 78 | elif mode == 1: |
| 79 | transform = [SphRandomRotate(p=1, img_size=img_size), |
| 80 | transforms.ToTensor(), |
| 81 | transforms.Normalize([0.5], [0.5])] |
| 82 | else: |
| 83 | raise NotImplementedError(f'Uncognized data augmentation mode') |
| 84 | else: |
| 85 | transform = [transforms.ToTensor(), |
| 86 | transforms.Resize(img_size, Image.BICUBIC), |
| 87 | transforms.Normalize([0.5], [0.5])] |
| 88 | self.transform = transforms.Compose(transform) |
| 89 | |
| 90 | def __call__(self, input): |
| 91 | output = self.transform(input) |
| 92 | return output |
| 93 | |
| 94 | |
| 95 | class SphRandomRotate: |