MCPcopy Create free account
hub / github.com/OpenTalker/StyleHEAT / ADAINDecoder

Class ADAINDecoder

models/styleheat/base_function.py:61–89  ·  view source on GitHub ↗

docstring for ADAINDecoder

Source from the content-addressed store, hash-verified

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

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected