MCPcopy Create free account
hub / github.com/drinkingcoder/FlowFormer-Official / BottleneckBlock

Class BottleneckBlock

core/extractor.py:60–116  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

58
59
60class BottleneckBlock(nn.Module):
61 def __init__(self, in_planes, planes, norm_fn='group', stride=1):
62 super(BottleneckBlock, self).__init__()
63
64 self.conv1 = nn.Conv2d(in_planes, planes//4, kernel_size=1, padding=0)
65 self.conv2 = nn.Conv2d(planes//4, planes//4, kernel_size=3, padding=1, stride=stride)
66 self.conv3 = nn.Conv2d(planes//4, planes, kernel_size=1, padding=0)
67 self.relu = nn.ReLU(inplace=True)
68
69 num_groups = planes // 8
70
71 if norm_fn == 'group':
72 self.norm1 = nn.GroupNorm(num_groups=num_groups, num_channels=planes//4)
73 self.norm2 = nn.GroupNorm(num_groups=num_groups, num_channels=planes//4)
74 self.norm3 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
75 if not stride == 1:
76 self.norm4 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
77
78 elif norm_fn == 'batch':
79 self.norm1 = nn.BatchNorm2d(planes//4)
80 self.norm2 = nn.BatchNorm2d(planes//4)
81 self.norm3 = nn.BatchNorm2d(planes)
82 if not stride == 1:
83 self.norm4 = nn.BatchNorm2d(planes)
84
85 elif norm_fn == 'instance':
86 self.norm1 = nn.InstanceNorm2d(planes//4)
87 self.norm2 = nn.InstanceNorm2d(planes//4)
88 self.norm3 = nn.InstanceNorm2d(planes)
89 if not stride == 1:
90 self.norm4 = nn.InstanceNorm2d(planes)
91
92 elif norm_fn == 'none':
93 self.norm1 = nn.Sequential()
94 self.norm2 = nn.Sequential()
95 self.norm3 = nn.Sequential()
96 if not stride == 1:
97 self.norm4 = nn.Sequential()
98
99 if stride == 1:
100 self.downsample = None
101
102 else:
103 self.downsample = nn.Sequential(
104 nn.Conv2d(in_planes, planes, kernel_size=1, stride=stride), self.norm4)
105
106
107 def forward(self, x):
108 y = x
109 y = self.relu(self.norm1(self.conv1(y)))
110 y = self.relu(self.norm2(self.conv2(y)))
111 y = self.relu(self.norm3(self.conv3(y)))
112
113 if self.downsample is not None:
114 x = self.downsample(x)
115
116 return self.relu(x+y)
117

Callers 1

_make_layerMethod · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected