MCPcopy Create free account
hub / github.com/ChunmingHe/WS-SAM / ETM

Class ETM

lib/Modules.py:160–194  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

158
159
160class ETM(nn.Module):
161 def __init__(self, in_channels, out_channels):
162 super(ETM, self).__init__()
163 self.relu = nn.ReLU(True)
164 self.branch0 = BasicConv2d(in_channels, out_channels, 1)
165 self.branch1 = nn.Sequential(
166 BasicConv2d(in_channels, out_channels, 1),
167 BasicConv2d(out_channels, out_channels, kernel_size=(1, 3), padding=(0, 1)),
168 BasicConv2d(out_channels, out_channels, kernel_size=(3, 1), padding=(1, 0)),
169 BasicConv2d(out_channels, out_channels, 3, padding=3, dilation=3)
170 )
171 self.branch2 = nn.Sequential(
172 BasicConv2d(in_channels, out_channels, 1),
173 BasicConv2d(out_channels, out_channels, kernel_size=(1, 5), padding=(0, 2)),
174 BasicConv2d(out_channels, out_channels, kernel_size=(5, 1), padding=(2, 0)),
175 BasicConv2d(out_channels, out_channels, 3, padding=5, dilation=5)
176 )
177 self.branch3 = nn.Sequential(
178 BasicConv2d(in_channels, out_channels, 1),
179 BasicConv2d(out_channels, out_channels, kernel_size=(1, 7), padding=(0, 3)),
180 BasicConv2d(out_channels, out_channels, kernel_size=(7, 1), padding=(3, 0)),
181 BasicConv2d(out_channels, out_channels, 3, padding=7, dilation=7)
182 )
183 self.conv_cat = BasicConv2d(4 * out_channels, out_channels, 3, padding=1)
184 self.conv_res = BasicConv2d(in_channels, out_channels, 1)
185
186 def forward(self, x):
187 x0 = self.branch0(x)
188 x1 = self.branch1(x)
189 x2 = self.branch2(x)
190 x3 = self.branch3(x)
191 x_cat = self.conv_cat(torch.cat((x0, x1, x2, x3), 1))
192
193 x = self.relu(x_cat + self.conv_res(x))
194 return x
195
196
197

Callers 7

__init__Method · 0.85
__init__Method · 0.85
__init__Method · 0.85
__init__Method · 0.85
__init__Method · 0.85
__init__Method · 0.85
__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected