| 185 | |
| 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 | |
| 213 | hidden = self.z2init(latent) |
| 214 | hidden = torch.split(hidden, self.hidden_size, dim=-1) |
| 215 | |
| 216 | return list(hidden) |
| 217 | |
| 218 | def forward(self, inputs, hidden, p): |
| 219 | # print(inputs.shape) |
| 220 | x_in = self.emb(inputs) |
| 221 | pos_enc = self.positional_encoder(p).to(inputs.device).detach() |
| 222 | x_in = x_in + pos_enc |
| 223 | |
| 224 | for i in range(self.n_layers): |
| 225 | hidden[i] = self.gru[i](x_in, hidden[i]) |
| 226 | h_in = hidden[i] |
| 227 | mu = self.mu_net(h_in) |
| 228 | logvar = self.logvar_net(h_in) |
| 229 | z = reparameterize(mu, logvar) |
| 230 | return z, mu, logvar, hidden |
| 231 | |
| 232 | class AttLayer(nn.Module): |
| 233 | def __init__(self, query_dim, key_dim, value_dim): |
nothing calls this directly
no outgoing calls
no test coverage detected