MCPcopy Create free account
hub / github.com/Xiaobin-Rong/gtcrn / Encoder

Class Encoder

gtcrn.py:228–244  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

226
227
228class Encoder(nn.Module):
229 def __init__(self):
230 super().__init__()
231 self.en_convs = nn.ModuleList([
232 ConvBlock(3*3, 16, (1,5), stride=(1,2), padding=(0,2), use_deconv=False, is_last=False),
233 ConvBlock(16, 16, (1,5), stride=(1,2), padding=(0,2), groups=2, use_deconv=False, is_last=False),
234 GTConvBlock(16, 16, (3,3), stride=(1,1), padding=(0,1), dilation=(1,1), use_deconv=False),
235 GTConvBlock(16, 16, (3,3), stride=(1,1), padding=(0,1), dilation=(2,1), use_deconv=False),
236 GTConvBlock(16, 16, (3,3), stride=(1,1), padding=(0,1), dilation=(5,1), use_deconv=False)
237 ])
238
239 def forward(self, x):
240 en_outs = []
241 for i in range(len(self.en_convs)):
242 x = self.en_convs[i](x)
243 en_outs.append(x)
244 return x, en_outs
245
246
247class Decoder(nn.Module):

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected