Differentiable Data Augmentation, intergrated resizing, shifting(ie, padding + cropping) and flipping. Input batch must be square images. Args: source_size(int): height of input images. target_size(int): height of output images. shift(int): maximum of allowd shifting si
| 34 | return len(self.dataset) |
| 35 | |
| 36 | class RandomTransform(torch.nn.Module): |
| 37 | """ Differentiable Data Augmentation, intergrated resizing, shifting(ie, padding + cropping) and flipping. Input batch must be square images. |
| 38 | |
| 39 | Args: |
| 40 | source_size(int): height of input images. |
| 41 | target_size(int): height of output images. |
| 42 | shift(int): maximum of allowd shifting size. |
| 43 | fliplr(bool): if flip horizonally |
| 44 | flipud(bool): if flip vertically |
| 45 | mode(string): the interpolation mode used in data augmentation. Default: bilinear. |
| 46 | align: the align mode used in data augmentation. Default: True. |
| 47 | |
| 48 | For more details, refers to https://discuss.pytorch.org/t/cropping-a-minibatch-of-images-each-image-a-bit-differently/12247/5 |
| 49 | """ |
| 50 | |
| 51 | def __init__(self, source_size, target_size, shift=8, fliplr=True, flipud=False, mode='bilinear', align=True): |
| 52 | """Args: source and target size.""" |
| 53 | super().__init__() |
| 54 | self.grid = self.build_grid(source_size, target_size) |
| 55 | self.delta = torch.linspace(0, 1, source_size)[shift] |
| 56 | self.fliplr = fliplr |
| 57 | self.flipud = flipud |
| 58 | self.mode = mode |
| 59 | self.align = True |
| 60 | |
| 61 | @staticmethod |
| 62 | def build_grid(source_size, target_size): |
| 63 | """https://discuss.pytorch.org/t/cropping-a-minibatch-of-images-each-image-a-bit-differently/12247/5.""" |
| 64 | k = float(target_size) / float(source_size) |
| 65 | direct = torch.linspace(-1, k, target_size).unsqueeze(0).repeat(target_size, 1).unsqueeze(-1) |
| 66 | full = torch.cat([direct, direct.transpose(1, 0)], dim=2).unsqueeze(0) |
| 67 | return full |
| 68 | |
| 69 | def random_crop_grid(self, x, randgen=None): |
| 70 | """https://discuss.pytorch.org/t/cropping-a-minibatch-of-images-each-image-a-bit-differently/12247/5.""" |
| 71 | grid = self.grid.repeat(x.size(0), 1, 1, 1).clone().detach() |
| 72 | grid = grid.to(device=x.device, dtype=x.dtype) |
| 73 | if randgen is None: |
| 74 | randgen = torch.rand(x.shape[0], 4, device=x.device, dtype=x.dtype) |
| 75 | |
| 76 | # Add random shifts by x |
| 77 | x_shift = (randgen[:, 0] - 0.5) * 2 * self.delta |
| 78 | grid[:, :, :, 0] = grid[:, :, :, 0] + x_shift.unsqueeze(-1).unsqueeze(-1).expand(-1, grid.size(1), grid.size(2)) |
| 79 | # Add random shifts by y |
| 80 | y_shift = (randgen[:, 1] - 0.5) * 2 * self.delta |
| 81 | grid[:, :, :, 1] = grid[:, :, :, 1] + y_shift.unsqueeze(-1).unsqueeze(-1).expand(-1, grid.size(1), grid.size(2)) |
| 82 | |
| 83 | if self.fliplr: |
| 84 | grid[randgen[:, 2] > 0.5, :, :, 0] *= -1 |
| 85 | if self.flipud: |
| 86 | grid[randgen[:, 3] > 0.5, :, :, 1] *= -1 |
| 87 | return grid |
| 88 | |
| 89 | |
| 90 | def forward(self, x, randgen=None): |
| 91 | # Make a random shift grid for each batch |
| 92 | grid_shifted = self.random_crop_grid(x, randgen) |
| 93 | # Sample using grid sample |