MCPcopy Create free account
hub / github.com/THUYimingLi/BackdoorBox / RandomTransform

Class RandomTransform

core/attacks/SleeperAgent.py:36–94  ·  view source on GitHub ↗

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

Source from the content-addressed store, hash-verified

34 return len(self.dataset)
35
36class 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

Callers 1

trainMethod · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected