MCPcopy Create free account
hub / github.com/Kitware/COAT / erase_patch

Class erase_patch

utils/mask.py:246–270  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

244
245
246class 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
272class mixup_patch:
273 def __init__(self, mask_size=2):

Callers 1

forwardMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected