MCPcopy Create free account
hub / github.com/ali-vilab/dreamtalk / ADAINDecoder

Class ADAINDecoder

generators/base_function.py:64–90  ·  view source on GitHub ↗

docstring for ADAINDecoder

Source from the content-addressed store, hash-verified

62 return out_list
63
64class ADAINDecoder(nn.Module):
65 """docstring for ADAINDecoder"""
66 def __init__(self, pose_nc, ngf, img_f, encoder_layers, decoder_layers, skip_connect=True,
67 nonlinearity=nn.LeakyReLU(), use_spect=False):
68
69 super(ADAINDecoder, self).__init__()
70 self.encoder_layers = encoder_layers
71 self.decoder_layers = decoder_layers
72 self.skip_connect = skip_connect
73 use_transpose = True
74
75 for i in range(encoder_layers-decoder_layers, encoder_layers)[::-1]:
76 in_channels = min(ngf * (2**(i+1)), img_f)
77 in_channels = in_channels*2 if i != (encoder_layers-1) and self.skip_connect else in_channels
78 out_channels = min(ngf * (2**i), img_f)
79 model = ADAINDecoderBlock(in_channels, out_channels, out_channels, pose_nc, use_transpose, nonlinearity, use_spect)
80 setattr(self, 'decoder' + str(i), model)
81
82 self.output_nc = out_channels*2 if self.skip_connect else out_channels
83
84 def forward(self, x, z):
85 out = x.pop() if self.skip_connect else x
86 for i in range(self.encoder_layers-self.decoder_layers, self.encoder_layers)[::-1]:
87 model = getattr(self, 'decoder' + str(i))
88 out = model(out, z)
89 out = torch.cat([out, x.pop()], 1) if self.skip_connect else out
90 return out
91
92class ADAINEncoderBlock(nn.Module):
93 def __init__(self, input_nc, output_nc, feature_nc, nonlinearity=nn.LeakyReLU(), use_spect=False):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected