AttMultiHeadAdd 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 (doe
(self, enc_hs_pad, enc_hs_len, dec_z, att_prev)
| 1054 | self.mask = None |
| 1055 | |
| 1056 | def forward(self, enc_hs_pad, enc_hs_len, dec_z, att_prev): |
| 1057 | """AttMultiHeadAdd forward |
| 1058 | |
| 1059 | :param torch.Tensor enc_hs_pad: padded encoder hidden state (B x T_max x D_enc) |
| 1060 | :param list enc_hs_len: padded encoder hidden state length (B) |
| 1061 | :param torch.Tensor dec_z: decoder hidden state (B x D_dec) |
| 1062 | :param torch.Tensor att_prev: dummy (does not use) |
| 1063 | :return: attention weighted encoder state (B, D_enc) |
| 1064 | :rtype: torch.Tensor |
| 1065 | :return: list of previous attention weight (B x T_max) * aheads |
| 1066 | :rtype: list |
| 1067 | """ |
| 1068 | |
| 1069 | batch = enc_hs_pad.size(0) |
| 1070 | # pre-compute all k and v outside the decoder loop |
| 1071 | if self.pre_compute_k is None or self.han_mode: |
| 1072 | self.enc_h = enc_hs_pad # utt x frame x hdim |
| 1073 | self.h_length = self.enc_h.size(1) |
| 1074 | # utt x frame x att_dim |
| 1075 | self.pre_compute_k = [self.mlp_k[h](self.enc_h) for h in six.moves.range(self.aheads)] |
| 1076 | |
| 1077 | if self.pre_compute_v is None or self.han_mode: |
| 1078 | self.enc_h = enc_hs_pad # utt x frame x hdim |
| 1079 | self.h_length = self.enc_h.size(1) |
| 1080 | # utt x frame x att_dim |
| 1081 | self.pre_compute_v = [self.mlp_v[h](self.enc_h) for h in six.moves.range(self.aheads)] |
| 1082 | |
| 1083 | if dec_z is None: |
| 1084 | dec_z = enc_hs_pad.new_zeros(batch, self.dunits) |
| 1085 | else: |
| 1086 | dec_z = dec_z.view(batch, self.dunits) |
| 1087 | |
| 1088 | c = [] |
| 1089 | w = [] |
| 1090 | for h in six.moves.range(self.aheads): |
| 1091 | e = self.gvec[h]( |
| 1092 | torch.tanh( |
| 1093 | self.pre_compute_k[h] + self.mlp_q[h](dec_z).view(batch, 1, self.att_dim_k) |
| 1094 | ) |
| 1095 | ).squeeze(2) |
| 1096 | |
| 1097 | # NOTE consider zero padding when compute w. |
| 1098 | if self.mask is None: |
| 1099 | self.mask = to_device(enc_hs_pad, make_pad_mask(enc_hs_len)) |
| 1100 | e.masked_fill_(self.mask, -float("inf")) |
| 1101 | w += [F.softmax(self.scaling * e, dim=1)] |
| 1102 | |
| 1103 | # weighted sum over flames |
| 1104 | # utt x hdim |
| 1105 | # NOTE use bmm instead of sum(*) |
| 1106 | c += [torch.sum(self.pre_compute_v[h] * w[h].view(batch, self.h_length, 1), dim=1)] |
| 1107 | |
| 1108 | # concat all of c |
| 1109 | c = self.mlp_o(torch.cat(c, dim=1)) |
| 1110 | |
| 1111 | return c, w |
| 1112 | |
| 1113 |
nothing calls this directly
no test coverage detected