| 244 | |
| 245 | |
| 246 | class erase_patch: |
| 247 | def __init__(self, mask_size=2): |
| 248 | self.mask_size = mask_size |
| 249 | |
| 250 | def __call__(self, features): |
| 251 | std, mean = torch.std_mean(features.detach()) |
| 252 | dim = features.shape[1] |
| 253 | if random.random() < 0.5: |
| 254 | y1_max = features.shape[3] - self.mask_size |
| 255 | num_masks = 1 |
| 256 | for i in range(num_masks): |
| 257 | mask_y1 = torch.randint(y1_max, (features.shape[0],)) |
| 258 | mask_y2 = mask_y1 + self.mask_size |
| 259 | for k in range(features.shape[0]): |
| 260 | features[k, :, :, mask_y1[k] : mask_y2[k]] = torch.normal(mean.repeat(dim,14,2), std.repeat(dim,14,2)) |
| 261 | else: |
| 262 | x1_max = features.shape[3] - self.mask_size |
| 263 | num_masks = 1 |
| 264 | for i in range(num_masks): |
| 265 | mask_x1 = torch.randint(x1_max, (features.shape[0],)) |
| 266 | mask_x2 = mask_x1 + self.mask_size |
| 267 | for k in range(features.shape[0]): |
| 268 | features[k, :, mask_x1[k] : mask_x2[k], :] = torch.normal(mean.repeat(dim,2,14), std.repeat(dim,2,14)) |
| 269 | |
| 270 | return features |
| 271 | |
| 272 | class mixup_patch: |
| 273 | def __init__(self, mask_size=2): |