| 230 | return z, mu, logvar, hidden |
| 231 | |
| 232 | class AttLayer(nn.Module): |
| 233 | def __init__(self, query_dim, key_dim, value_dim): |
| 234 | super(AttLayer, self).__init__() |
| 235 | self.W_q = nn.Linear(query_dim, value_dim) |
| 236 | self.W_k = nn.Linear(key_dim, value_dim, bias=False) |
| 237 | self.W_v = nn.Linear(key_dim, value_dim) |
| 238 | |
| 239 | self.softmax = nn.Softmax(dim=1) |
| 240 | self.dim = value_dim |
| 241 | |
| 242 | self.W_q.apply(init_weight) |
| 243 | self.W_k.apply(init_weight) |
| 244 | self.W_v.apply(init_weight) |
| 245 | |
| 246 | def forward(self, query, key_mat): |
| 247 | ''' |
| 248 | query (batch, query_dim) |
| 249 | key (batch, seq_len, key_dim) |
| 250 | ''' |
| 251 | # print(query.shape) |
| 252 | query_vec = self.W_q(query).unsqueeze(-1) # (batch, value_dim, 1) |
| 253 | val_set = self.W_v(key_mat) # (batch, seq_len, value_dim) |
| 254 | key_set = self.W_k(key_mat) # (batch, seq_len, value_dim) |
| 255 | |
| 256 | weights = torch.matmul(key_set, query_vec) / np.sqrt(self.dim) |
| 257 | |
| 258 | co_weights = self.softmax(weights) # (batch, seq_len, 1) |
| 259 | values = val_set * co_weights # (batch, seq_len, value_dim) |
| 260 | pred = values.sum(dim=1) # (batch, value_dim) |
| 261 | return pred, co_weights |
| 262 | |
| 263 | def short_cut(self, querys, keys): |
| 264 | return self.W_q(querys), self.W_k(keys) |
| 265 | |
| 266 | |
| 267 | class TextEncoderBiGRU(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected