MCPcopy Create free account
hub / github.com/modelscope/FunASR / forward

Method forward

funasr/models/language_model/rnn/attentions.py:1056–1111  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 3

to_deviceFunction · 0.90
make_pad_maskFunction · 0.90
softmaxMethod · 0.45

Tested by

no test coverage detected