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

Class ADAINDecoderBlock

generators/base_function.py:111–148  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

109 return x
110
111class ADAINDecoderBlock(nn.Module):
112 def __init__(self, input_nc, output_nc, hidden_nc, feature_nc, use_transpose=True, nonlinearity=nn.LeakyReLU(), use_spect=False):
113 super(ADAINDecoderBlock, self).__init__()
114 # Attributes
115 self.actvn = nonlinearity
116 hidden_nc = min(input_nc, output_nc) if hidden_nc is None else hidden_nc
117
118 kwargs_fine = {'kernel_size':3, 'stride':1, 'padding':1}
119 if use_transpose:
120 kwargs_up = {'kernel_size':3, 'stride':2, 'padding':1, 'output_padding':1}
121 else:
122 kwargs_up = {'kernel_size':3, 'stride':1, 'padding':1}
123
124 # create conv layers
125 self.conv_0 = spectral_norm(nn.Conv2d(input_nc, hidden_nc, **kwargs_fine), use_spect)
126 if use_transpose:
127 self.conv_1 = spectral_norm(nn.ConvTranspose2d(hidden_nc, output_nc, **kwargs_up), use_spect)
128 self.conv_s = spectral_norm(nn.ConvTranspose2d(input_nc, output_nc, **kwargs_up), use_spect)
129 else:
130 self.conv_1 = nn.Sequential(spectral_norm(nn.Conv2d(hidden_nc, output_nc, **kwargs_up), use_spect),
131 nn.Upsample(scale_factor=2))
132 self.conv_s = nn.Sequential(spectral_norm(nn.Conv2d(input_nc, output_nc, **kwargs_up), use_spect),
133 nn.Upsample(scale_factor=2))
134 # define normalization layers
135 self.norm_0 = ADAIN(input_nc, feature_nc)
136 self.norm_1 = ADAIN(hidden_nc, feature_nc)
137 self.norm_s = ADAIN(input_nc, feature_nc)
138
139 def forward(self, x, z):
140 x_s = self.shortcut(x, z)
141 dx = self.conv_0(self.actvn(self.norm_0(x, z)))
142 dx = self.conv_1(self.actvn(self.norm_1(dx, z)))
143 out = x_s + dx
144 return out
145
146 def shortcut(self, x, z):
147 x_s = self.conv_s(self.actvn(self.norm_s(x, z)))
148 return x_s
149
150
151def spectral_norm(module, use_spect=True):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected