(self, text_size, input_size, output_size, hidden_size, n_layers)
| 186 | |
| 187 | class TextDecoder(nn.Module): |
| 188 | def __init__(self, text_size, input_size, output_size, hidden_size, n_layers): |
| 189 | super(TextDecoder, self).__init__() |
| 190 | self.input_size = input_size |
| 191 | self.output_size = output_size |
| 192 | self.hidden_size = hidden_size |
| 193 | self.n_layers = n_layers |
| 194 | self.emb = nn.Sequential( |
| 195 | nn.Linear(input_size, hidden_size), |
| 196 | nn.LayerNorm(hidden_size), |
| 197 | nn.LeakyReLU(0.2, inplace=True)) |
| 198 | |
| 199 | self.gru = nn.ModuleList([nn.GRUCell(hidden_size, hidden_size) for i in range(self.n_layers)]) |
| 200 | self.z2init = nn.Linear(text_size, hidden_size * n_layers) |
| 201 | self.positional_encoder = PositionalEncoding(hidden_size) |
| 202 | |
| 203 | self.mu_net = nn.Linear(hidden_size, output_size) |
| 204 | self.logvar_net = nn.Linear(hidden_size, output_size) |
| 205 | |
| 206 | self.emb.apply(init_weight) |
| 207 | self.z2init.apply(init_weight) |
| 208 | self.mu_net.apply(init_weight) |
| 209 | self.logvar_net.apply(init_weight) |
| 210 | |
| 211 | def get_init_hidden(self, latent): |
| 212 |
nothing calls this directly
no test coverage detected