AttAdd forward :param torch.Tensor enc_hs_pad: padded encoder hidden state (B x T_max x D_enc) :param list enc_hs_len: padded encoder hidden state length (B) :param torch.Tensor dec_z: decoder hidden state (B x D_dec) :param torch.Tensor att_prev: dummy (does not use
(self, enc_hs_pad, enc_hs_len, dec_z, att_prev, scaling=2.0)
| 212 | self.mask = None |
| 213 | |
| 214 | def forward(self, enc_hs_pad, enc_hs_len, dec_z, att_prev, scaling=2.0): |
| 215 | """AttAdd forward |
| 216 | |
| 217 | :param torch.Tensor enc_hs_pad: padded encoder hidden state (B x T_max x D_enc) |
| 218 | :param list enc_hs_len: padded encoder hidden state length (B) |
| 219 | :param torch.Tensor dec_z: decoder hidden state (B x D_dec) |
| 220 | :param torch.Tensor att_prev: dummy (does not use) |
| 221 | :param float scaling: scaling parameter before applying softmax |
| 222 | :return: attention weighted encoder state (B, D_enc) |
| 223 | :rtype: torch.Tensor |
| 224 | :return: previous attention weights (B x T_max) |
| 225 | :rtype: torch.Tensor |
| 226 | """ |
| 227 | |
| 228 | batch = len(enc_hs_pad) |
| 229 | # pre-compute all h outside the decoder loop |
| 230 | if self.pre_compute_enc_h is None or self.han_mode: |
| 231 | self.enc_h = enc_hs_pad # utt x frame x hdim |
| 232 | self.h_length = self.enc_h.size(1) |
| 233 | # utt x frame x att_dim |
| 234 | self.pre_compute_enc_h = self.mlp_enc(self.enc_h) |
| 235 | |
| 236 | if dec_z is None: |
| 237 | dec_z = enc_hs_pad.new_zeros(batch, self.dunits) |
| 238 | else: |
| 239 | dec_z = dec_z.view(batch, self.dunits) |
| 240 | |
| 241 | # dec_z_tiled: utt x frame x att_dim |
| 242 | dec_z_tiled = self.mlp_dec(dec_z).view(batch, 1, self.att_dim) |
| 243 | |
| 244 | # dot with gvec |
| 245 | # utt x frame x att_dim -> utt x frame |
| 246 | e = self.gvec(torch.tanh(self.pre_compute_enc_h + dec_z_tiled)).squeeze(2) |
| 247 | |
| 248 | # NOTE consider zero padding when compute w. |
| 249 | if self.mask is None: |
| 250 | self.mask = to_device(enc_hs_pad, make_pad_mask(enc_hs_len)) |
| 251 | e.masked_fill_(self.mask, -float("inf")) |
| 252 | w = F.softmax(scaling * e, dim=1) |
| 253 | |
| 254 | # weighted sum over flames |
| 255 | # utt x hdim |
| 256 | # NOTE use bmm instead of sum(*) |
| 257 | c = torch.sum(self.enc_h * w.view(batch, self.h_length, 1), dim=1) |
| 258 | |
| 259 | return c, w |
| 260 | |
| 261 | |
| 262 | class AttLoc(torch.nn.Module): |
nothing calls this directly
no test coverage detected