| 133 | |
| 134 | |
| 135 | class exchange_patch: |
| 136 | def __init__(self, shape='stripe', mask_size=2, mode='random_direct'): |
| 137 | self.shape = shape |
| 138 | self.mask_size = mask_size |
| 139 | self.mode = mode |
| 140 | |
| 141 | def __call__(self, features): |
| 142 | # Stripe mask |
| 143 | if self.shape == 'stripe': |
| 144 | if self.mode == 'horizontal': |
| 145 | features = self.xpatch_hstripe(features, self.mask_size) |
| 146 | elif self.mode == 'vertical': |
| 147 | features = self.xpatch_vstripe(features, self.mask_size) |
| 148 | elif self.mode == 'random_direction': |
| 149 | if random.random() < 0.5: |
| 150 | features = self.xpatch_hstripe(features, self.mask_size) |
| 151 | else: |
| 152 | features = self.xpatch_vstripe(features, self.mask_size) |
| 153 | else: |
| 154 | raise Exception("Unknown stripe mask mode name") |
| 155 | # Square mask |
| 156 | elif self.shape == 'square': |
| 157 | if self.mode == 'random_size': |
| 158 | self.mask_size = 4 if random.random() < 0.5 else 5 |
| 159 | features = self.xpatch_square(features, self.mask_size) |
| 160 | # Random stripe/square mask |
| 161 | elif self.shape == 'random': |
| 162 | random_num = random.random() |
| 163 | if random_num < 0.25: |
| 164 | features = self.xpatch_hstripe(features, 2) |
| 165 | elif random_num < 0.5 and random_num >= 0.25: |
| 166 | features = self.xpatch_vstripe(features, 2) |
| 167 | elif random_num < 0.75 and random_num >= 0.5: |
| 168 | features = self.xpatch_square(features, 4) |
| 169 | else: |
| 170 | features = self.xpatch_square(features, 5) |
| 171 | else: |
| 172 | raise Exception("Unknown mask shape name") |
| 173 | |
| 174 | return features |
| 175 | |
| 176 | def xpatch_hstripe(self, features, mask_size): |
| 177 | """ |
| 178 | """ |
| 179 | # horizontal stripe |
| 180 | y1_max = features.shape[3] - mask_size |
| 181 | num_masks = 1 |
| 182 | for i in range(num_masks): |
| 183 | mask_y1 = torch.randint(y1_max, (1,)) |
| 184 | mask_y2 = mask_y1 + mask_size |
| 185 | new_idx = torch.randperm(features.shape[0]) |
| 186 | features[:, :, :, mask_y1 : mask_y2] = features[new_idx, :, :, mask_y1 : mask_y2] |
| 187 | return features |
| 188 | |
| 189 | |
| 190 | def xpatch_vstripe(self, features, mask_size): |
| 191 | """ |
| 192 | """ |