| 5 | |
| 6 | |
| 7 | class DropBlock(nn.Module): |
| 8 | def __init__(self, block_size): |
| 9 | super(DropBlock, self).__init__() |
| 10 | |
| 11 | self.block_size = block_size |
| 12 | |
| 13 | def forward(self, x, gamma): |
| 14 | # shape: (bsize, channels, height, width) |
| 15 | |
| 16 | if self.training: |
| 17 | batch_size, channels, height, width = x.shape |
| 18 | bernoulli = Bernoulli(gamma) |
| 19 | mask = bernoulli.sample((batch_size, channels, height - (self.block_size - 1), width - (self.block_size - 1))) |
| 20 | if torch.cuda.is_available(): |
| 21 | mask = mask.cuda() |
| 22 | block_mask = self._compute_block_mask(mask) |
| 23 | countM = block_mask.size()[0] * block_mask.size()[1] * block_mask.size()[2] * block_mask.size()[3] |
| 24 | count_ones = block_mask.sum() |
| 25 | |
| 26 | return block_mask * x * (countM / count_ones) |
| 27 | else: |
| 28 | return x |
| 29 | |
| 30 | def _compute_block_mask(self, mask): |
| 31 | left_padding = int((self.block_size-1) / 2) |
| 32 | right_padding = int(self.block_size / 2) |
| 33 | |
| 34 | batch_size, channels, height, width = mask.shape |
| 35 | non_zero_idxs = mask.nonzero() |
| 36 | nr_blocks = non_zero_idxs.shape[0] |
| 37 | |
| 38 | offsets = torch.stack( |
| 39 | [ |
| 40 | torch.arange(self.block_size).view(-1, 1).expand(self.block_size, self.block_size).reshape(-1), # - left_padding, |
| 41 | torch.arange(self.block_size).repeat(self.block_size), #- left_padding |
| 42 | ] |
| 43 | ).t() |
| 44 | offsets = torch.cat((torch.zeros(self.block_size**2, 2).long(), offsets.long()), 1) |
| 45 | if torch.cuda.is_available(): |
| 46 | offsets = offsets.cuda() |
| 47 | |
| 48 | if nr_blocks > 0: |
| 49 | non_zero_idxs = non_zero_idxs.repeat(self.block_size ** 2, 1) |
| 50 | offsets = offsets.repeat(nr_blocks, 1).view(-1, 4) |
| 51 | offsets = offsets.long() |
| 52 | |
| 53 | block_idxs = non_zero_idxs + offsets |
| 54 | #block_idxs += left_padding |
| 55 | padded_mask = F.pad(mask, (left_padding, right_padding, left_padding, right_padding)) |
| 56 | padded_mask[block_idxs[:, 0], block_idxs[:, 1], block_idxs[:, 2], block_idxs[:, 3]] = 1. |
| 57 | else: |
| 58 | padded_mask = F.pad(mask, (left_padding, right_padding, left_padding, right_padding)) |
| 59 | |
| 60 | block_mask = 1 - padded_mask#[:height, :width] |
| 61 | return block_mask |