(self, inputs, hidden, p)
| 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 test coverage detected