MCPcopy Create free account
hub / github.com/DanielShalam/BPA / DropBlock

Class DropBlock

models/dropblock.py:7–61  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

5
6
7class 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

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected