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

Class exchange_patch

utils/mask.py:135–218  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

133
134
135class 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 """

Callers 2

__init__Method · 0.90
forwardMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected